device_test.go 6.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252
  1. package device
  2. import (
  3. "encoding/json"
  4. "os"
  5. "path/filepath"
  6. "strings"
  7. "testing"
  8. )
  9. func TestIndexDevices(t *testing.T) {
  10. tests := []struct {
  11. name string
  12. devices []Device
  13. want []string
  14. }{
  15. {
  16. name: "empty",
  17. devices: nil,
  18. want: nil,
  19. },
  20. {
  21. name: "single",
  22. devices: []Device{{Path: "/dev/hidraw7", VendorID: 0x36b0, ProductID: 0x309f}},
  23. want: []string{"/dev/hidraw7"},
  24. },
  25. {
  26. name: "sorted by vendor then product then path",
  27. devices: []Device{
  28. {Path: "/dev/hidraw9", VendorID: 0x9999, ProductID: 0x0001},
  29. {Path: "/dev/hidraw5", VendorID: 0x1111, ProductID: 0x0002},
  30. {Path: "/dev/hidraw2", VendorID: 0x1111, ProductID: 0x0001},
  31. },
  32. want: []string{"/dev/hidraw2", "/dev/hidraw5", "/dev/hidraw9"},
  33. },
  34. {
  35. name: "identical vendor and product sorted by path",
  36. devices: []Device{
  37. {Path: "/dev/hidrawB", VendorID: 0x36b0, ProductID: 0x309f},
  38. {Path: "/dev/hidrawA", VendorID: 0x36b0, ProductID: 0x309f},
  39. },
  40. want: []string{"/dev/hidrawA", "/dev/hidrawB"},
  41. },
  42. }
  43. for _, tt := range tests {
  44. t.Run(tt.name, func(t *testing.T) {
  45. got := indexDevices(tt.devices)
  46. if len(got) != len(tt.want) {
  47. t.Fatalf("indexDevices() returned %d devices, want %d", len(got), len(tt.want))
  48. }
  49. for i, wantPath := range tt.want {
  50. if got[i].Path != wantPath {
  51. t.Errorf("indexDevices()[%d].Path = %q, want %q", i, got[i].Path, wantPath)
  52. }
  53. if got[i].Index != i+1 {
  54. t.Errorf("indexDevices()[%d].Index = %d, want %d", i, got[i].Index, i+1)
  55. }
  56. }
  57. })
  58. }
  59. }
  60. func TestIndexDevicesDoesNotMutateInput(t *testing.T) {
  61. input := []Device{
  62. {Path: "/dev/hidraw9", VendorID: 0x9999, ProductID: 0x0001},
  63. {Path: "/dev/hidraw2", VendorID: 0x1111, ProductID: 0x0001},
  64. }
  65. indexDevices(input)
  66. if input[0].Path != "/dev/hidraw9" {
  67. t.Errorf("indexDevices() mutated input: input[0].Path = %q, want %q", input[0].Path, "/dev/hidraw9")
  68. }
  69. if input[0].Index != 0 {
  70. t.Errorf("indexDevices() mutated input: input[0].Index = %d, want 0", input[0].Index)
  71. }
  72. }
  73. func TestDeviceJSONFieldNames(t *testing.T) {
  74. dev := Device{
  75. Index: 1,
  76. Path: "/dev/hidraw7",
  77. VendorID: 0x36b0,
  78. ProductID: 0x309f,
  79. Name: "Wobkey Impact 80",
  80. Known: true,
  81. }
  82. data, err := json.Marshal(dev)
  83. if err != nil {
  84. t.Fatalf("Marshal() returned error: %v", err)
  85. }
  86. var got map[string]any
  87. if err := json.Unmarshal(data, &got); err != nil {
  88. t.Fatalf("Unmarshal() returned error: %v", err)
  89. }
  90. for _, field := range []string{"index", "path", "vendorId", "productId", "name", "known"} {
  91. if _, ok := got[field]; !ok {
  92. t.Errorf("Device JSON = %s, missing field %q", data, field)
  93. }
  94. }
  95. if _, ok := got["supported"]; ok {
  96. t.Errorf("Device JSON = %s, must not contain the misleading field \"supported\"", data)
  97. }
  98. }
  99. func TestDeviceJSONOmitsUnnamedName(t *testing.T) {
  100. data, err := json.Marshal(Device{Index: 2, Path: "/dev/hidraw9", Known: false})
  101. if err != nil {
  102. t.Fatalf("Marshal() returned error: %v", err)
  103. }
  104. if strings.Contains(string(data), `"name"`) {
  105. t.Errorf("Device JSON = %s, want the name omitted for an unknown keyboard", data)
  106. }
  107. }
  108. func TestKeyboardIgnoresUnknownFields(t *testing.T) {
  109. tmpDir := t.TempDir()
  110. tmpFile := filepath.Join(tmpDir, "keyboards.json")
  111. // Entries may still carry fields the tool does not use; loading must
  112. // succeed and yield only the fields the code actually reads.
  113. testData := `[
  114. {"name":"Impact 80","vendorId":14000,"productId":12447,
  115. "protocol":"via","viaVersion":3,"ledLayout":"rgblight","futureField":"ignored"}
  116. ]`
  117. if err := os.WriteFile(tmpFile, []byte(testData), 0644); err != nil {
  118. t.Fatalf("failed to write test file: %v", err)
  119. }
  120. origFind := findKeyboardsJSON
  121. findKeyboardsJSON = func() (string, error) { return tmpFile, nil }
  122. defer func() { findKeyboardsJSON = origFind }()
  123. keyboards, err := LoadKeyboards()
  124. if err != nil {
  125. t.Fatalf("LoadKeyboards() returned error: %v", err)
  126. }
  127. if len(keyboards) != 1 {
  128. t.Fatalf("LoadKeyboards() returned %d keyboards, want 1", len(keyboards))
  129. }
  130. if keyboards[0].Name != "Impact 80" {
  131. t.Errorf("Name = %q, want %q", keyboards[0].Name, "Impact 80")
  132. }
  133. }
  134. func TestKeyboardJSONOmitsRemovedFields(t *testing.T) {
  135. data, err := json.Marshal(Keyboard{Name: "Impact 80", VendorID: 0x36b0, ProductID: 0x309f})
  136. if err != nil {
  137. t.Fatalf("Marshal() returned error: %v", err)
  138. }
  139. for _, dead := range []string{"protocol", "viaVersion", "ledLayout"} {
  140. if strings.Contains(string(data), dead) {
  141. t.Errorf("Keyboard JSON = %s, want no %q field", data, dead)
  142. }
  143. }
  144. }
  145. func TestLoadKeyboards(t *testing.T) {
  146. tmpDir := t.TempDir()
  147. tmpFile := filepath.Join(tmpDir, "keyboards.json")
  148. testData := `[
  149. {"name": "Test Keyboard 1", "vendorId": 22222, "productId": 1, "protocol": "via"},
  150. {"name": "Test Keyboard 2", "vendorId": 33333, "productId": 2, "protocol": "viapro"}
  151. ]`
  152. if err := os.WriteFile(tmpFile, []byte(testData), 0644); err != nil {
  153. t.Fatalf("failed to write test file: %v", err)
  154. }
  155. origFind := findKeyboardsJSON
  156. findKeyboardsJSON = func() (string, error) {
  157. return tmpFile, nil
  158. }
  159. defer func() { findKeyboardsJSON = origFind }()
  160. keyboards, err := LoadKeyboards()
  161. if err != nil {
  162. t.Fatalf("LoadKeyboards() returned error: %v", err)
  163. }
  164. if len(keyboards) != 2 {
  165. t.Errorf("LoadKeyboards() returned %d keyboards, want 2", len(keyboards))
  166. }
  167. if keyboards[0].Name != "Test Keyboard 1" {
  168. t.Errorf("First keyboard name = %q, want %q", keyboards[0].Name, "Test Keyboard 1")
  169. }
  170. if keyboards[0].VendorID != 22222 {
  171. t.Errorf("First keyboard VID = %d, want %d", keyboards[0].VendorID, 22222)
  172. }
  173. if keyboards[1].ProductID != 2 {
  174. t.Errorf("Second keyboard PID = %d, want %d", keyboards[1].ProductID, 2)
  175. }
  176. }
  177. func TestLoadKeyboardsInvalid(t *testing.T) {
  178. tmpDir := t.TempDir()
  179. tmpFile := filepath.Join(tmpDir, "keyboards.json")
  180. invalidData := `[this is not valid json`
  181. if err := os.WriteFile(tmpFile, []byte(invalidData), 0644); err != nil {
  182. t.Fatalf("failed to write test file: %v", err)
  183. }
  184. origFind := findKeyboardsJSON
  185. findKeyboardsJSON = func() (string, error) {
  186. return tmpFile, nil
  187. }
  188. defer func() { findKeyboardsJSON = origFind }()
  189. _, err := LoadKeyboards()
  190. if err == nil {
  191. t.Error("LoadKeyboards() expected error for invalid JSON, got nil")
  192. }
  193. }
  194. func TestLoadKeyboardsEmptyFile(t *testing.T) {
  195. tmpDir := t.TempDir()
  196. tmpFile := filepath.Join(tmpDir, "keyboards.json")
  197. if err := os.WriteFile(tmpFile, []byte("[]"), 0644); err != nil {
  198. t.Fatalf("failed to write test file: %v", err)
  199. }
  200. origFind := findKeyboardsJSON
  201. findKeyboardsJSON = func() (string, error) {
  202. return tmpFile, nil
  203. }
  204. defer func() { findKeyboardsJSON = origFind }()
  205. keyboards, err := LoadKeyboards()
  206. if err != nil {
  207. t.Fatalf("LoadKeyboards() returned error: %v", err)
  208. }
  209. if len(keyboards) != 0 {
  210. t.Errorf("LoadKeyboards() returned %d keyboards, want 0", len(keyboards))
  211. }
  212. }