Browse Source

Apply effects per lighting zone

Paul Klumpp 2 weeks ago
parent
commit
fd347a4348
3 changed files with 157 additions and 12 deletions
  1. 31 12
      cmd/wobkey/rgb/effect.go
  2. 97 0
      cmd/wobkey/rgb/effect_test.go
  3. 29 0
      internal/via/protocol_test.go

+ 31 - 12
cmd/wobkey/rgb/effect.go

@@ -3,38 +3,57 @@ package rgb
 import (
 	"fmt"
 	"os"
+	"strings"
 
 	"github.com/spf13/cobra"
-	"github.com/wobkey/rgb/internal/rgb"
+	intrgb "github.com/wobkey/rgb/internal/rgb"
 	"github.com/wobkey/rgb/internal/via"
 )
 
+func resolveEffectTargets(name string, zones []intrgb.Zone) ([]intrgb.EffectTarget, []intrgb.Zone, error) {
+	if name == "static" {
+		name = "solid"
+	}
+	return intrgb.ResolveEffect(name, zones)
+}
+
 func NewEffectCmd() *cobra.Command {
 	return &cobra.Command{
 		Use:   "effect <name>",
 		Short: "Set RGB effect",
 		Long:  "Set the RGB lighting effect on the connected keyboard.",
 		Args:  cobra.ExactArgs(1),
-		Run: func(cmd *cobra.Command, args []string) {
-			proto, _, err := OpenDevice()
+		RunE: func(_ *cobra.Command, args []string) error {
+			zones, err := selectedZones()
 			if err != nil {
-				fmt.Fprintf(os.Stderr, "Error: %v\n", err)
-				os.Exit(1)
+				return err
+			}
+			targets, skipped, err := resolveEffectTargets(args[0], zones)
+			if err != nil {
+				return err
+			}
+			if len(skipped) > 0 {
+				skippedNames := make([]string, len(skipped))
+				for i, zone := range skipped {
+					skippedNames[i] = string(zone)
+				}
+				fmt.Fprintf(os.Stderr, "Warning: effect not supported on zone(s): %s\n", strings.Join(skippedNames, ", "))
 			}
-			defer proto.Close()
 
-			e, err := rgb.ParseImpact80Effect(args[0])
+			proto, _, err := OpenDevice()
 			if err != nil {
-				fmt.Fprintf(os.Stderr, "Error: %v\n", err)
-				os.Exit(1)
+				return err
 			}
+			defer proto.Close()
 
-			if err := proto.SetValue(via.RGBLight, uint8(rgb.EffectID), e); err != nil {
-				fmt.Fprintf(os.Stderr, "Error setting effect: %v\n", err)
-				os.Exit(1)
+			for _, target := range targets {
+				if err := proto.SetValue(via.LEDType(target.Zone.Channel()), uint8(intrgb.EffectID), target.ID); err != nil {
+					return fmt.Errorf("set effect on %s: %w", target.Zone, err)
+				}
 			}
 
 			fmt.Printf("Effect set to %q\n", args[0])
+			return nil
 		},
 	}
 }

+ 97 - 0
cmd/wobkey/rgb/effect_test.go

@@ -0,0 +1,97 @@
+package rgb
+
+import (
+	"reflect"
+	"testing"
+
+	intrgb "github.com/wobkey/rgb/internal/rgb"
+)
+
+func TestResolveEffectTargetsBacklightOnly(t *testing.T) {
+	targets, skipped, err := resolveEffectTargets("rainbow_moving_chevron", intrgb.AllZones())
+	if err != nil {
+		t.Fatalf("resolveEffectTargets() unexpected error: %v", err)
+	}
+	wantTargets := []intrgb.EffectTarget{{Zone: intrgb.ZoneBacklight, ID: 17}}
+	if !reflect.DeepEqual(targets, wantTargets) {
+		t.Errorf("resolveEffectTargets() targets = %v, want %v", targets, wantTargets)
+	}
+	wantSkipped := []intrgb.Zone{intrgb.ZoneLogo, intrgb.ZoneSide}
+	if !reflect.DeepEqual(skipped, wantSkipped) {
+		t.Errorf("resolveEffectTargets() skipped = %v, want %v", skipped, wantSkipped)
+	}
+}
+
+func TestResolveEffectTargetsBreathing(t *testing.T) {
+	targets, skipped, err := resolveEffectTargets("breathing", intrgb.AllZones())
+	if err != nil {
+		t.Fatalf("resolveEffectTargets() unexpected error: %v", err)
+	}
+	wantTargets := []intrgb.EffectTarget{
+		{Zone: intrgb.ZoneLogo, ID: 4},
+		{Zone: intrgb.ZoneBacklight, ID: 5},
+		{Zone: intrgb.ZoneSide, ID: 4},
+	}
+	if !reflect.DeepEqual(targets, wantTargets) {
+		t.Errorf("resolveEffectTargets() targets = %v, want %v", targets, wantTargets)
+	}
+	if len(skipped) != 0 {
+		t.Errorf("resolveEffectTargets() skipped = %v, want none", skipped)
+	}
+}
+
+func TestResolveEffectTargetsExplicitUnsupported(t *testing.T) {
+	cases := []struct {
+		name  string
+		zones []intrgb.Zone
+	}{
+		{"rainbow_moving_chevron", []intrgb.Zone{intrgb.ZoneLogo}},
+		{"rainbow_wave", []intrgb.Zone{intrgb.ZoneBacklight}},
+		{"splash", []intrgb.Zone{intrgb.ZoneSide}},
+	}
+	for _, tc := range cases {
+		t.Run(tc.name+string(tc.zones[0]), func(t *testing.T) {
+			targets, skipped, err := resolveEffectTargets(tc.name, tc.zones)
+			if err == nil {
+				t.Fatal("resolveEffectTargets() expected error, got nil")
+			}
+			if len(targets) != 0 {
+				t.Errorf("resolveEffectTargets() targets = %v, want none", targets)
+			}
+			if len(skipped) != 0 {
+				t.Errorf("resolveEffectTargets() skipped = %v, want none", skipped)
+			}
+		})
+	}
+}
+
+func TestResolveEffectTargetsUnknownName(t *testing.T) {
+	targets, skipped, err := resolveEffectTargets("not_an_effect", intrgb.AllZones())
+	if err == nil {
+		t.Fatal("resolveEffectTargets() expected error, got nil")
+	}
+	if len(targets) != 0 {
+		t.Errorf("resolveEffectTargets() targets = %v, want none", targets)
+	}
+	if len(skipped) != 0 {
+		t.Errorf("resolveEffectTargets() skipped = %v, want none", skipped)
+	}
+}
+
+func TestResolveEffectTargetsStaticCompatibility(t *testing.T) {
+	targets, skipped, err := resolveEffectTargets("static", intrgb.AllZones())
+	if err != nil {
+		t.Fatalf("resolveEffectTargets() unexpected error: %v", err)
+	}
+	wantTargets := []intrgb.EffectTarget{
+		{Zone: intrgb.ZoneLogo, ID: 5},
+		{Zone: intrgb.ZoneBacklight, ID: 1},
+		{Zone: intrgb.ZoneSide, ID: 5},
+	}
+	if !reflect.DeepEqual(targets, wantTargets) {
+		t.Errorf("resolveEffectTargets() targets = %v, want %v", targets, wantTargets)
+	}
+	if len(skipped) != 0 {
+		t.Errorf("resolveEffectTargets() skipped = %v, want none", skipped)
+	}
+}

+ 29 - 0
internal/via/protocol_test.go

@@ -197,3 +197,32 @@ func TestGetValueRejectsUnhandledResponse(t *testing.T) {
 		t.Fatal("GetValue() error = nil, want unhandled response error")
 	}
 }
+
+func TestSetValueWritesEachImpact80EffectChannel(t *testing.T) {
+	transport := &fakeTransport{}
+	protocol := Protocol{handle: transport}
+	cases := []struct {
+		channel LEDType
+		effect  byte
+	}{
+		{RGBLight, 4},
+		{RGBMatrix, 5},
+		{SideLight, 4},
+	}
+
+	for _, tc := range cases {
+		if err := protocol.SetValue(tc.channel, 0x02, tc.effect); err != nil {
+			t.Fatalf("SetValue(%d) error = %v", tc.channel, err)
+		}
+	}
+
+	if len(transport.reports) != len(cases) {
+		t.Fatalf("SendReport() calls = %d, want %d", len(transport.reports), len(cases))
+	}
+	for i, tc := range cases {
+		report := transport.reports[i]
+		if report[0] != 0x07 || report[1] != byte(tc.channel) || report[2] != 0x02 || report[3] != tc.effect {
+			t.Errorf("effect report %d = %v, want channel 0x%02x effect %d", i, report, tc.channel, tc.effect)
+		}
+	}
+}