Bläddra i källkod

detect which VIA lighting channels a keyboard has

Paul Klumpp 1 vecka sedan
förälder
incheckning
2558f78cf5
4 ändrade filer med 248 tillägg och 145 borttagningar
  1. 73 0
      internal/via/channel.go
  2. 132 0
      internal/via/channel_test.go
  3. 18 62
      internal/via/protocol.go
  4. 25 83
      internal/via/protocol_test.go

+ 73 - 0
internal/via/channel.go

@@ -0,0 +1,73 @@
+package via
+
+import (
+	"errors"
+	"fmt"
+)
+
+// Channel is a VIA lighting channel. The values are QMK's id_qmk_*_channel
+// constants from quantum/via.h.
+type Channel uint8
+
+const (
+	ChannelBacklight Channel = 1
+	ChannelRgblight  Channel = 2
+	ChannelRgbMatrix Channel = 3
+	ChannelAudio     Channel = 4
+	ChannelLedMatrix Channel = 5
+)
+
+// AssignedChannelMax is the highest channel QMK currently assigns. The probe
+// asks beyond it on purpose, see probeChannelMax.
+const AssignedChannelMax = ChannelLedMatrix
+
+// probeChannelMax is the highest channel the probe asks about. The QMK enum
+// stops at 5; the extra reads cost one round trip each and let a future QMK
+// that claims a higher channel be found without a code change.
+const probeChannelMax = 15
+
+// probeValueID is the brightness value ID. It is one byte and valid on every
+// lighting subsystem, so one request per channel finds any kind of channel.
+const probeValueID = 0x01
+
+// Subsystem returns the QMK name of the channel's lighting subsystem. The name
+// follows from the channel number, so no board file stores it and it is always
+// available for a channel the keyboard has.
+func (c Channel) Subsystem() string {
+	switch c {
+	case ChannelBacklight:
+		return "backlight"
+	case ChannelRgblight:
+		return "rgblight"
+	case ChannelRgbMatrix:
+		return "rgb_matrix"
+	case ChannelAudio:
+		return "audio"
+	case ChannelLedMatrix:
+		return "led_matrix"
+	default:
+		return ""
+	}
+}
+
+// DetectChannels asks the keyboard which lighting channels it has, in ascending
+// order.
+//
+// The 0xFF answer is the discriminator, not the payload: QMK rejects a channel
+// it does not compile in with id_unhandled, whereas an unknown value ID on a
+// known channel is mirrored back with a zero payload. A read that fails for any
+// other reason aborts the probe, because a channel list missing entries is
+// indistinguishable from a complete one and would silently under-report.
+func (p *Protocol) DetectChannels() ([]Channel, error) {
+	var present []Channel
+	for c := Channel(1); c <= probeChannelMax; c++ {
+		if _, err := p.GetValue(c, probeValueID); err != nil {
+			if errors.Is(err, errUnhandled) {
+				continue
+			}
+			return nil, fmt.Errorf("probe channel %d: %w", c, err)
+		}
+		present = append(present, c)
+	}
+	return present, nil
+}

+ 132 - 0
internal/via/channel_test.go

@@ -0,0 +1,132 @@
+package via
+
+import (
+	"errors"
+	"strings"
+	"testing"
+)
+
+// probeQueue scripts one answer per channel the probe visits, so a test can say
+// which channels the keyboard has.
+func probeQueue(present ...Channel) [][]byte {
+	queue := make([][]byte, 0, probeChannelMax)
+	for c := 1; c <= probeChannelMax; c++ {
+		buf := make([]byte, 32)
+		if containsChannel(present, Channel(c)) {
+			buf[0] = byte(CustomGet)
+			buf[3] = 160
+		} else {
+			buf[0] = byte(Unhandled)
+		}
+		buf[1] = byte(c)
+		buf[2] = probeValueID
+		queue = append(queue, buf)
+	}
+	return queue
+}
+
+func containsChannel(channels []Channel, want Channel) bool {
+	for _, c := range channels {
+		if c == want {
+			return true
+		}
+	}
+	return false
+}
+
+func TestDetectChannelsFindsTheAnsweredChannels(t *testing.T) {
+	transport := &fakeTransport{queue: probeQueue(ChannelRgblight, ChannelRgbMatrix, ChannelAudio)}
+	protocol := Protocol{handle: transport}
+
+	got, err := protocol.DetectChannels()
+	if err != nil {
+		t.Fatalf("DetectChannels() error = %v", err)
+	}
+
+	want := []Channel{ChannelRgblight, ChannelRgbMatrix, ChannelAudio}
+	if len(got) != len(want) {
+		t.Fatalf("DetectChannels() = %v, want %v", got, want)
+	}
+	for i := range want {
+		if got[i] != want[i] {
+			t.Errorf("DetectChannels()[%d] = %d, want %d", i, got[i], want[i])
+		}
+	}
+}
+
+func TestDetectChannelsAsksBrightnessOnEveryChannelFromOneToFifteen(t *testing.T) {
+	transport := &fakeTransport{queue: probeQueue()}
+	protocol := Protocol{handle: transport}
+
+	if _, err := protocol.DetectChannels(); err != nil {
+		t.Fatalf("DetectChannels() error = %v", err)
+	}
+
+	if len(transport.reports) != probeChannelMax {
+		t.Fatalf("requests = %d, want %d", len(transport.reports), probeChannelMax)
+	}
+	for i, report := range transport.reports {
+		wantChannel := byte(i + 1)
+		if report[0] != byte(CustomGet) {
+			t.Errorf("request %d command = 0x%02x, want 0x%02x", i, report[0], CustomGet)
+		}
+		if report[1] != wantChannel {
+			t.Errorf("request %d channel = %d, want %d", i, report[1], wantChannel)
+		}
+		if report[2] != probeValueID {
+			t.Errorf("request %d value = 0x%02x, want 0x%02x (brightness)", i, report[2], probeValueID)
+		}
+	}
+}
+
+// A channel list missing entries is indistinguishable from a complete one, so a
+// transport failure has to abort the probe instead of shrinking it.
+func TestDetectChannelsFailsOnATransportErrorMidProbe(t *testing.T) {
+	transport := &fakeTransport{
+		queue:        probeQueue(ChannelRgblight, ChannelRgbMatrix, ChannelAudio),
+		readErr:      errors.New("read: interrupted system call"),
+		readErrAfter: 3,
+	}
+	protocol := Protocol{handle: transport}
+
+	_, err := protocol.DetectChannels()
+	if err == nil {
+		t.Fatal("DetectChannels() expected an error, got nil")
+	}
+	if !strings.Contains(err.Error(), "channel 4") {
+		t.Errorf("error = %q, want it to name the channel that failed", err)
+	}
+}
+
+func TestDetectChannelsReportsNoChannelsWhenNoneArePresent(t *testing.T) {
+	transport := &fakeTransport{queue: probeQueue()}
+	protocol := Protocol{handle: transport}
+
+	got, err := protocol.DetectChannels()
+	if err != nil {
+		t.Fatalf("DetectChannels() error = %v", err)
+	}
+	if len(got) != 0 {
+		t.Errorf("DetectChannels() = %v, want empty", got)
+	}
+}
+
+func TestChannelSubsystemNames(t *testing.T) {
+	tests := []struct {
+		channel Channel
+		want    string
+	}{
+		{ChannelBacklight, "backlight"},
+		{ChannelRgblight, "rgblight"},
+		{ChannelRgbMatrix, "rgb_matrix"},
+		{ChannelAudio, "audio"},
+		{ChannelLedMatrix, "led_matrix"},
+		{Channel(9), ""},
+	}
+
+	for _, tt := range tests {
+		if got := tt.channel.Subsystem(); got != tt.want {
+			t.Errorf("Channel(%d).Subsystem() = %q, want %q", tt.channel, got, tt.want)
+		}
+	}
+}

+ 18 - 62
internal/via/protocol.go

@@ -1,6 +1,7 @@
 package via
 
 import (
+	"errors"
 	"fmt"
 
 	"netdome.biz/paul/qmk-rgb/internal/device"
@@ -45,75 +46,30 @@ const (
 	Unhandled Message = 0xff
 )
 
-// LEDType defines QMK LED subsystem types.
-type LEDType uint8
+// errUnhandled reports a channel or value ID the firmware does not implement.
+// It is a sentinel so a caller can tell "this does not exist here" from a
+// transport failure, which the two are otherwise indistinguishable in.
+var errUnhandled = errors.New("unhandled response")
 
-const (
-	RGBLight  LEDType = 0x02
-	RGBMatrix LEDType = 0x03
-	SideLight LEDType = 0x04
-)
-
-// SetValue sends a set value command for the given LED type and parameter.
-func (p *Protocol) SetValue(ledType LEDType, param uint8, value uint8) error {
+// SetValue sends a set value command for the given channel and parameter.
+func (p *Protocol) SetValue(ch Channel, param uint8, value uint8) error {
 	report := make([]byte, 32)
 	report[0] = byte(CustomSet)
-	report[1] = byte(ledType)
+	report[1] = byte(ch)
 	report[2] = param
 	report[3] = value
 
 	if _, err := p.handle.SendReport(0x00, report); err != nil {
 		return err
 	}
-	_, err := p.readResponse(CustomSet, ledType, param)
+	_, err := p.readResponse(CustomSet, byte(ch), param)
 	return err
 }
 
-func (p *Protocol) DisableLighting() error {
-	for _, channel := range []LEDType{RGBLight, RGBMatrix, SideLight} {
-		if err := p.SetValue(channel, 0x02, 0x00); err != nil {
-			return err
-		}
-		if err := p.SetValue(channel, 0x01, 0x00); err != nil {
-			return err
-		}
-	}
-	return nil
-}
-
-func (p *Protocol) EnableLighting() error {
-	channels := []struct {
-		channel LEDType
-		effect  uint8
-	}{
-		{RGBLight, 0x04},
-		{RGBMatrix, 0x05},
-		{SideLight, 0x04},
-	}
-	for _, channel := range channels {
-		if err := p.SetValue(channel.channel, 0x02, channel.effect); err != nil {
-			return err
-		}
-		if err := p.SetValue(channel.channel, 0x01, 160); err != nil {
-			return err
-		}
-	}
-	return nil
-}
-
-func (p *Protocol) SetLightingColor(hue uint8, saturation uint8) error {
-	for _, channel := range []LEDType{RGBLight, RGBMatrix, SideLight} {
-		if err := p.SetColor(channel, hue, saturation); err != nil {
-			return err
-		}
-	}
-	return nil
-}
-
-func (p *Protocol) SetColor(ledType LEDType, hue uint8, saturation uint8) error {
+func (p *Protocol) SetColor(ch Channel, hue uint8, saturation uint8) error {
 	report := make([]byte, 32)
 	report[0] = byte(CustomSet)
-	report[1] = byte(ledType)
+	report[1] = byte(ch)
 	report[2] = 0x04
 	report[3] = hue
 	report[4] = saturation
@@ -121,15 +77,15 @@ func (p *Protocol) SetColor(ledType LEDType, hue uint8, saturation uint8) error
 	if _, err := p.handle.SendReport(0x00, report); err != nil {
 		return err
 	}
-	_, err := p.readResponse(CustomSet, ledType, 0x04)
+	_, err := p.readResponse(CustomSet, byte(ch), 0x04)
 	return err
 }
 
 // GetValue sends a get value request and reads the response.
-func (p *Protocol) GetValue(ledType LEDType, param uint8) ([]byte, error) {
+func (p *Protocol) GetValue(ch Channel, param uint8) ([]byte, error) {
 	report := make([]byte, 32)
 	report[0] = byte(CustomGet)
-	report[1] = byte(ledType)
+	report[1] = byte(ch)
 	report[2] = param
 
 	_, err := p.handle.SendReport(0x00, report)
@@ -137,7 +93,7 @@ func (p *Protocol) GetValue(ledType LEDType, param uint8) ([]byte, error) {
 		return nil, fmt.Errorf("send get request: %w", err)
 	}
 
-	buf, err := p.readResponse(CustomGet, ledType, param)
+	buf, err := p.readResponse(CustomGet, byte(ch), param)
 	if err != nil {
 		return nil, err
 	}
@@ -149,7 +105,7 @@ func (p *Protocol) GetValue(ledType LEDType, param uint8) ([]byte, error) {
 	return append([]byte(nil), buf[3:3+valueSize]...), nil
 }
 
-func (p *Protocol) readResponse(command Message, ledType LEDType, param uint8) ([]byte, error) {
+func (p *Protocol) readResponse(command Message, ch, param byte) ([]byte, error) {
 	buf := make([]byte, 32)
 	n, err := p.handle.Read(buf)
 	if err != nil {
@@ -159,12 +115,12 @@ func (p *Protocol) readResponse(command Message, ledType LEDType, param uint8) (
 		return nil, fmt.Errorf("short response: got %d bytes, want 32", n)
 	}
 	if buf[0] == byte(Unhandled) {
-		return nil, fmt.Errorf("unhandled response")
+		return nil, errUnhandled
 	}
 	if buf[0] != byte(command) {
 		return nil, fmt.Errorf("unexpected response command: 0x%02x", buf[0])
 	}
-	if buf[1] != byte(ledType) || buf[2] != param {
+	if buf[1] != ch || buf[2] != param {
 		return nil, fmt.Errorf("response value mismatch: got channel 0x%02x value 0x%02x", buf[1], buf[2])
 	}
 	return buf, nil

+ 25 - 83
internal/via/protocol_test.go

@@ -6,10 +6,13 @@ import (
 )
 
 type fakeTransport struct {
-	reportIDs []byte
-	reports   [][]byte
-	response  []byte
-	readCalls int
+	reportIDs    []byte
+	reports      [][]byte
+	response     []byte
+	queue        [][]byte
+	readCalls    int
+	readErr      error
+	readErrAfter int
 }
 
 func (f *fakeTransport) SendReport(reportID byte, report []byte) (int, error) {
@@ -20,6 +23,14 @@ func (f *fakeTransport) SendReport(reportID byte, report []byte) (int, error) {
 
 func (f *fakeTransport) Read(buf []byte) (int, error) {
 	f.readCalls++
+	if f.readErr != nil && f.readCalls > f.readErrAfter {
+		return 0, f.readErr
+	}
+	if len(f.queue) > 0 {
+		next := f.queue[0]
+		f.queue = f.queue[1:]
+		return copy(buf, next), nil
+	}
 	if len(f.response) > 0 {
 		return copy(buf, f.response), nil
 	}
@@ -35,7 +46,7 @@ func TestSetValueUsesQMKRGBLightPayload(t *testing.T) {
 	transport := &fakeTransport{}
 	protocol := Protocol{handle: transport}
 
-	if err := protocol.SetValue(RGBLight, 0x02, 0x00); err != nil {
+	if err := protocol.SetValue(ChannelRgblight, 0x02, 0x00); err != nil {
 		t.Fatalf("SetValue() error = %v", err)
 	}
 
@@ -65,7 +76,7 @@ func TestSetValueConsumesQMKResponse(t *testing.T) {
 	transport := &fakeTransport{response: response}
 	protocol := Protocol{handle: transport}
 
-	if err := protocol.SetValue(RGBLight, 0x02, 0x00); err != nil {
+	if err := protocol.SetValue(ChannelRgblight, 0x02, 0x00); err != nil {
 		t.Fatalf("SetValue() error = %v", err)
 	}
 	if transport.readCalls != 1 {
@@ -73,31 +84,11 @@ func TestSetValueConsumesQMKResponse(t *testing.T) {
 	}
 }
 
-func TestSetLightingColorUsesAllImpact80Channels(t *testing.T) {
-	transport := &fakeTransport{}
-	protocol := Protocol{handle: transport}
-
-	if err := protocol.SetLightingColor(0x55, 0xff); err != nil {
-		t.Fatalf("SetLightingColor() error = %v", err)
-	}
-
-	wantChannels := []byte{0x02, 0x03, 0x04}
-	if len(transport.reports) != len(wantChannels) {
-		t.Fatalf("SendReport() calls = %d, want %d", len(transport.reports), len(wantChannels))
-	}
-	for i, channel := range wantChannels {
-		report := transport.reports[i]
-		if report[0] != 0x07 || report[1] != channel || report[2] != 0x04 || report[3] != 0x55 || report[4] != 0xff {
-			t.Errorf("color report %d = %v, want channel 0x%02x hue 0x55 saturation 0xff", i, report, channel)
-		}
-	}
-}
-
 func TestSetColorUsesQMKColorValue(t *testing.T) {
 	transport := &fakeTransport{}
 	protocol := Protocol{handle: transport}
 
-	if err := protocol.SetColor(RGBLight, 0x2a, 0x80); err != nil {
+	if err := protocol.SetColor(ChannelRgblight, 0x2a, 0x80); err != nil {
 		t.Fatalf("SetColor() error = %v", err)
 	}
 
@@ -113,55 +104,6 @@ func TestSetColorUsesQMKColorValue(t *testing.T) {
 	}
 }
 
-func TestDisableLightingUsesAllImpact80Channels(t *testing.T) {
-	transport := &fakeTransport{}
-	protocol := Protocol{handle: transport}
-
-	if err := protocol.DisableLighting(); err != nil {
-		t.Fatalf("DisableLighting() error = %v", err)
-	}
-
-	if len(transport.reports) != 6 {
-		t.Fatalf("SendReport() calls = %d, want 6", len(transport.reports))
-	}
-	wantChannels := []byte{0x02, 0x03, 0x04}
-	for i, channel := range wantChannels {
-		effect := transport.reports[i*2]
-		brightness := transport.reports[i*2+1]
-		if effect[1] != channel || effect[2] != 0x02 || effect[3] != 0x00 {
-			t.Errorf("effect report %d = %v, want channel 0x%02x effect 0", i, effect, channel)
-		}
-		if brightness[1] != channel || brightness[2] != 0x01 || brightness[3] != 0x00 {
-			t.Errorf("brightness report %d = %v, want channel 0x%02x brightness 0", i, brightness, channel)
-		}
-	}
-}
-
-func TestEnableLightingUsesAllImpact80Channels(t *testing.T) {
-	transport := &fakeTransport{}
-	protocol := Protocol{handle: transport}
-
-	if err := protocol.EnableLighting(); err != nil {
-		t.Fatalf("EnableLighting() error = %v", err)
-	}
-
-	if len(transport.reports) != 6 {
-		t.Fatalf("SendReport() calls = %d, want 6", len(transport.reports))
-	}
-	wantChannels := []byte{0x02, 0x03, 0x04}
-	wantEffects := []byte{0x04, 0x05, 0x04}
-	for i, channel := range wantChannels {
-		effect := transport.reports[i*2]
-		brightness := transport.reports[i*2+1]
-		if effect[1] != channel || effect[2] != 0x02 || effect[3] != wantEffects[i] {
-			t.Errorf("effect report %d = %v, want channel 0x%02x effect %d", i, effect, channel, wantEffects[i])
-		}
-		if brightness[1] != channel || brightness[2] != 0x01 || brightness[3] != 160 {
-			t.Errorf("brightness report %d = %v, want channel 0x%02x brightness 160", i, brightness, channel)
-		}
-	}
-}
-
 func TestGetValueReturnsQMKValueData(t *testing.T) {
 	response := make([]byte, 32)
 	response[0] = 0x08
@@ -171,7 +113,7 @@ func TestGetValueReturnsQMKValueData(t *testing.T) {
 	transport := &fakeTransport{response: response}
 	protocol := Protocol{handle: transport}
 
-	got, err := protocol.GetValue(RGBLight, 0x01)
+	got, err := protocol.GetValue(ChannelRgblight, 0x01)
 	if err != nil {
 		t.Fatalf("GetValue() error = %v", err)
 	}
@@ -193,7 +135,7 @@ func TestGetValueRejectsUnhandledResponse(t *testing.T) {
 	response[0] = 0xff
 	protocol := Protocol{handle: &fakeTransport{response: response}}
 
-	if _, err := protocol.GetValue(RGBLight, 0x01); err == nil {
+	if _, err := protocol.GetValue(ChannelRgblight, 0x01); err == nil {
 		t.Fatal("GetValue() error = nil, want unhandled response error")
 	}
 }
@@ -202,12 +144,12 @@ func TestSetValueWritesEachImpact80EffectChannel(t *testing.T) {
 	transport := &fakeTransport{}
 	protocol := Protocol{handle: transport}
 	cases := []struct {
-		channel LEDType
+		channel Channel
 		effect  byte
 	}{
-		{RGBLight, 4},
-		{RGBMatrix, 5},
-		{SideLight, 4},
+		{ChannelRgblight, 4},
+		{ChannelRgbMatrix, 5},
+		{ChannelAudio, 4},
 	}
 
 	for _, tc := range cases {
@@ -237,7 +179,7 @@ func TestGetValueReturnsTwoByteColor(t *testing.T) {
 	transport := &fakeTransport{response: response}
 	protocol := Protocol{handle: transport}
 
-	got, err := protocol.GetValue(RGBMatrix, 0x04)
+	got, err := protocol.GetValue(ChannelRgbMatrix, 0x04)
 	if err != nil {
 		t.Fatalf("GetValue() error = %v", err)
 	}