agi.cnn_test.go 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218
  1. package agi
  2. import (
  3. "encoding/json"
  4. "net/http"
  5. "net/http/httptest"
  6. "strings"
  7. "testing"
  8. "github.com/robertkrimen/otto"
  9. "imuslab.com/arozos/mod/agi/static"
  10. user "imuslab.com/arozos/mod/user"
  11. )
  12. // ─── pure helpers ─────────────────────────────────────────────────────────────
  13. func TestCNNMaskToken(t *testing.T) {
  14. cases := map[string]string{
  15. "": "",
  16. "abc": "•••",
  17. "cxn-1234567890": "••••7890",
  18. }
  19. for in, want := range cases {
  20. if got := cnnMaskToken(in); got != want {
  21. t.Errorf("cnnMaskToken(%q) = %q, want %q", in, got, want)
  22. }
  23. }
  24. }
  25. func TestCNNIsImageExt(t *testing.T) {
  26. if !cnnIsImageExt(".png") || !cnnIsImageExt(".jpeg") || !cnnIsImageExt(".webp") {
  27. t.Error("expected image extensions to be detected")
  28. }
  29. if cnnIsImageExt(".txt") {
  30. t.Error(".txt should not be an image")
  31. }
  32. }
  33. func TestParseCNNOptions(t *testing.T) {
  34. if opt := parseCNNOptions(""); opt.Model != "" {
  35. t.Errorf("empty string should yield zero options")
  36. }
  37. if opt := parseCNNOptions("undefined"); opt.Model != "" {
  38. t.Errorf("'undefined' should yield zero options")
  39. }
  40. if opt := parseCNNOptions("null"); opt.Model != "" {
  41. t.Errorf("'null' should yield zero options")
  42. }
  43. opt := parseCNNOptions(`{"model":"yolo11n","score_threshold":0.3,"render":true,"top_k":5}`)
  44. if opt.Model != "yolo11n" || !opt.Render || opt.TopK != 5 {
  45. t.Errorf("unexpected parse: %+v", opt)
  46. }
  47. if opt.ScoreThreshold == nil || *opt.ScoreThreshold != 0.3 {
  48. t.Errorf("score_threshold not parsed: %+v", opt)
  49. }
  50. }
  51. func TestParseCNNComparisonOptions(t *testing.T) {
  52. opt := parseCNNComparisonOptions(`{"threshold":0.5,"a_cropped":true}`)
  53. if opt.Threshold == nil || *opt.Threshold != 0.5 || !opt.ACropped {
  54. t.Errorf("unexpected parse: %+v", opt)
  55. }
  56. if opt := parseCNNComparisonOptions(""); opt.Threshold != nil {
  57. t.Errorf("empty string should yield zero options")
  58. }
  59. }
  60. // ─── persistence ──────────────────────────────────────────────────────────────
  61. func TestCNNConfigDefaultsAndPersistence(t *testing.T) {
  62. g := dbGateway(t)
  63. sysdb := g.Option.UserHandler.GetDatabase()
  64. sysdb.NewTable(cnnDBTable)
  65. cfg := g.getCNNConfig()
  66. if cfg.TimeoutSeconds != cnnDefaultTimeoutSeconds {
  67. t.Errorf("expected default timeout %d, got %d", cnnDefaultTimeoutSeconds, cfg.TimeoutSeconds)
  68. }
  69. sysdb.Write(cnnDBTable, "config", CNNServerConfig{Endpoint: "http://localhost:8080", Token: "tok", TimeoutSeconds: 30})
  70. cfg = g.getCNNConfig()
  71. if cfg.Endpoint != "http://localhost:8080" || cfg.Token != "tok" || cfg.TimeoutSeconds != 30 {
  72. t.Errorf("unexpected config after write: %+v", cfg)
  73. }
  74. }
  75. func TestCNNClientRequiresEndpoint(t *testing.T) {
  76. g := dbGateway(t)
  77. sysdb := g.Option.UserHandler.GetDatabase()
  78. sysdb.NewTable(cnnDBTable)
  79. if _, err := g.cnnClient(); err == nil {
  80. t.Fatal("expected an error when endpoint is not configured")
  81. }
  82. }
  83. // ─── VM injection ─────────────────────────────────────────────────────────────
  84. // TestCNNFunctionsInjected verifies every cnn.* binding is exposed as a
  85. // function after injection. The file-reading bindings (classify, detect, ...)
  86. // are checked for existence only here, mirroring how aimodel.chatWithFile is
  87. // checked for the aimodel lib (agi.aimodel_test.go) - actually invoking them
  88. // needs a fully wired virtual filesystem + user permission set that isn't
  89. // modelled anywhere in this test suite.
  90. func TestCNNFunctionsInjected(t *testing.T) {
  91. g := dbGateway(t)
  92. sysdb := g.Option.UserHandler.GetDatabase()
  93. sysdb.NewTable(cnnDBTable)
  94. vm := otto.New()
  95. payload := &static.AgiLibInjectionPayload{VM: vm, User: &user.User{Username: "tester"}}
  96. g.injectCNNFunctions(payload)
  97. methods := []string{
  98. "classify", "detect", "segment", "pose", "oriented",
  99. "faceDetect", "faceLandmarks", "faceEmbedding", "faceAttributes",
  100. "faceCompare", "analyze", "job", "models", "health",
  101. }
  102. for _, method := range methods {
  103. val, err := vm.Run(`typeof cnn.` + method)
  104. if err != nil {
  105. t.Fatalf("evaluating cnn.%s: %v", method, err)
  106. }
  107. if s, _ := val.ToString(); s != "function" {
  108. t.Errorf("cnn.%s should be a function, got %q", method, s)
  109. }
  110. }
  111. }
  112. // TestCNNHealthAndModelsRoundTrip exercises the full native-func -> JSON ->
  113. // JS-shim round trip for the file-free bindings against a mock CXNNAIO server.
  114. func TestCNNHealthAndModelsRoundTrip(t *testing.T) {
  115. srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  116. switch r.URL.Path {
  117. case "/v1/health":
  118. json.NewEncoder(w).Encode(map[string]any{"status": "ok", "version": "0.1.0", "models_loaded": 12, "uptime_s": 100})
  119. case "/v1/models":
  120. json.NewEncoder(w).Encode(map[string]any{"object": "list", "data": []map[string]any{{"id": "yolo11n", "object": "model", "task": "detection"}}})
  121. default:
  122. http.NotFound(w, r)
  123. }
  124. }))
  125. defer srv.Close()
  126. g := dbGateway(t)
  127. sysdb := g.Option.UserHandler.GetDatabase()
  128. sysdb.NewTable(cnnDBTable)
  129. sysdb.Write(cnnDBTable, "config", CNNServerConfig{Endpoint: srv.URL, TimeoutSeconds: 5})
  130. vm := otto.New()
  131. g.injectCNNFunctions(&static.AgiLibInjectionPayload{VM: vm, User: &user.User{Username: "tester"}})
  132. val, err := vm.Run(`cnn.health().status`)
  133. if err != nil {
  134. t.Fatalf("cnn.health() errored: %v", err)
  135. }
  136. if s, _ := val.ToString(); s != "ok" {
  137. t.Errorf("expected status ok, got %q", s)
  138. }
  139. val, err = vm.Run(`cnn.models().data[0].id`)
  140. if err != nil {
  141. t.Fatalf("cnn.models() errored: %v", err)
  142. }
  143. if s, _ := val.ToString(); s != "yolo11n" {
  144. t.Errorf("expected yolo11n, got %q", s)
  145. }
  146. }
  147. // TestCNNJobPoll exercises the async job-poll binding end-to-end.
  148. func TestCNNJobPoll(t *testing.T) {
  149. srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  150. if r.URL.Path != "/v1/jobs/job-1" {
  151. t.Errorf("unexpected path: %s", r.URL.Path)
  152. return
  153. }
  154. json.NewEncoder(w).Encode(map[string]any{
  155. "id": "job-1", "object": "job", "status": "succeeded",
  156. "result": map[string]any{"object": "image.detection", "data": []any{}},
  157. })
  158. }))
  159. defer srv.Close()
  160. g := dbGateway(t)
  161. sysdb := g.Option.UserHandler.GetDatabase()
  162. sysdb.NewTable(cnnDBTable)
  163. sysdb.Write(cnnDBTable, "config", CNNServerConfig{Endpoint: srv.URL, TimeoutSeconds: 5})
  164. vm := otto.New()
  165. g.injectCNNFunctions(&static.AgiLibInjectionPayload{VM: vm, User: &user.User{Username: "tester"}})
  166. val, err := vm.Run(`cnn.job("job-1").status`)
  167. if err != nil {
  168. t.Fatalf("cnn.job() errored: %v", err)
  169. }
  170. if s, _ := val.ToString(); s != "succeeded" {
  171. t.Errorf("expected succeeded, got %q", s)
  172. }
  173. }
  174. // TestCNNHealthErrorsWhenUnconfigured checks the CNNError surfaces cleanly
  175. // when no endpoint has been saved yet.
  176. func TestCNNHealthErrorsWhenUnconfigured(t *testing.T) {
  177. g := dbGateway(t)
  178. sysdb := g.Option.UserHandler.GetDatabase()
  179. sysdb.NewTable(cnnDBTable)
  180. vm := otto.New()
  181. g.injectCNNFunctions(&static.AgiLibInjectionPayload{VM: vm, User: &user.User{Username: "tester"}})
  182. _, err := vm.Run(`cnn.health()`)
  183. if err == nil {
  184. t.Fatal("expected an error when CNN server is not configured")
  185. }
  186. if !strings.Contains(err.Error(), "not configured") {
  187. t.Errorf("expected a 'not configured' error, got: %v", err)
  188. }
  189. }