From ddd1aae5d13fa460b96b609901f1a25a62033244 Mon Sep 17 00:00:00 2001 From: jim800121chen Date: Wed, 22 Jul 2026 19:27:02 +0800 Subject: [PATCH] =?UTF-8?q?feat(server):=20=E6=8A=8A=20model=20metadata=20?= =?UTF-8?q?=E5=82=B3=E9=80=B2=E6=8E=A8=E8=AB=96=E9=8F=88=E8=B7=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit models.json 宣告的 taskType/labels/inputSize 原本在 flash 時被丟棄 (只傳 modelPath),導致 bridge 只能靠檔名猜測模型類型與尺寸。 - FlashOptions 帶 TaskType/Labels/InputWidth/InputHeight - 抽出 buildLoadModelCommand,四處 load_model 呼叫點(初次 + 三條 retry 路徑)統一走它,並加測試釘住呼叫點數量與「不得有手寫 payload」 —— 讓漏改 retry 路徑在結構上不可能發生 - ClassResult 加 ClassIndex(不加 omitempty,index 0 是合法值) 用 FlashOptions struct 而非裸參數,未來擴充欄位不需再動 interface 與所有 test fake。 Co-Authored-By: Claude Opus 4.8 (1M context) --- .../api/handlers/firmware_handler_test.go | 4 +- .../server/internal/device/manager_test.go | 4 +- .../server/internal/driver/interface.go | 52 ++- .../internal/driver/kneron/kl720_driver.go | 108 +++++- .../driver/kneron/load_model_payload_test.go | 313 ++++++++++++++++++ local-tool/server/internal/flash/service.go | 40 ++- .../internal/flash/service_metadata_test.go | 190 +++++++++++ 7 files changed, 690 insertions(+), 21 deletions(-) create mode 100644 local-tool/server/internal/driver/kneron/load_model_payload_test.go create mode 100644 local-tool/server/internal/flash/service_metadata_test.go diff --git a/local-tool/server/internal/api/handlers/firmware_handler_test.go b/local-tool/server/internal/api/handlers/firmware_handler_test.go index fa76487..0204e4e 100644 --- a/local-tool/server/internal/api/handlers/firmware_handler_test.go +++ b/local-tool/server/internal/api/handlers/firmware_handler_test.go @@ -125,7 +125,9 @@ func (f *fakeDriver) Info() driver.DeviceInfo { return f.in func (f *fakeDriver) Connect() error { return nil } func (f *fakeDriver) Disconnect() error { return nil } func (f *fakeDriver) IsConnected() bool { return false } -func (f *fakeDriver) Flash(_ string, _ chan<- driver.FlashProgress) error { return nil } +func (f *fakeDriver) Flash(_ string, _ driver.FlashOptions, _ chan<- driver.FlashProgress) error { + return nil +} func (f *fakeDriver) StartInference() error { return nil } func (f *fakeDriver) StopInference() error { return nil } func (f *fakeDriver) ReadInference() (*driver.InferenceResult, error) { diff --git a/local-tool/server/internal/device/manager_test.go b/local-tool/server/internal/device/manager_test.go index f0be750..962c297 100644 --- a/local-tool/server/internal/device/manager_test.go +++ b/local-tool/server/internal/device/manager_test.go @@ -15,7 +15,9 @@ func (d *testDriver) Info() driver.DeviceInfo { r func (d *testDriver) Connect() error { d.connected = true; d.info.Status = driver.StatusConnected; return nil } func (d *testDriver) Disconnect() error { d.connected = false; d.info.Status = driver.StatusDisconnected; return nil } func (d *testDriver) IsConnected() bool { return d.connected } -func (d *testDriver) Flash(_ string, _ chan<- driver.FlashProgress) error { return nil } +func (d *testDriver) Flash(_ string, _ driver.FlashOptions, _ chan<- driver.FlashProgress) error { + return nil +} func (d *testDriver) StartInference() error { return nil } func (d *testDriver) StopInference() error { return nil } func (d *testDriver) ReadInference() (*driver.InferenceResult, error) { return nil, nil } diff --git a/local-tool/server/internal/driver/interface.go b/local-tool/server/internal/driver/interface.go index 1a1ea9f..2b0c409 100644 --- a/local-tool/server/internal/driver/interface.go +++ b/local-tool/server/internal/driver/interface.go @@ -7,7 +7,7 @@ type DeviceDriver interface { Connect() error Disconnect() error IsConnected() bool - Flash(modelPath string, progressCh chan<- FlashProgress) error + Flash(modelPath string, opts FlashOptions, progressCh chan<- FlashProgress) error StartInference() error StopInference() error ReadInference() (*InferenceResult, error) @@ -39,6 +39,53 @@ const ( StatusDisconnected DeviceStatus = "disconnected" ) +// FlashOptions 帶入 model metadata,供 driver 在 load model 時傳給硬體 bridge。 +// +// 為什麼用 struct 而不是多帶兩個參數:載入模型需要的 metadata 之後還會長 +// (如前處理色彩格式、top-K),用 struct 之後新增欄位不必再改 interface 簽章 +// 與所有 test fake。 +// +// 兩個欄位都是 optional —— 空值代表「未指定」,bridge 端會 fallback 到既有的 +// model id / 檔名 heuristics(維持既有 detection 行為不變)。 +type FlashOptions struct { + // TaskType 為 models.json 宣告的推論類型("classification" / + // "object_detection")。bridge 端有指定就不再用檔名猜測。 + TaskType string + // Labels 是 class index → 顯示名稱的對應表,純顯示層用途、非推論必要輸入。 + // 沒帶時 classification 輸出原始 enum(class_N)、detection 沿用 COCO。 + Labels []string + // InputWidth / InputHeight 是 models.json / metadata.json 宣告的模型輸入 + // 尺寸。 + // + // ⚠️ 這是**最後手段**,不是可信來源:bridge 端會優先向 SDK 問模型自己 + // 宣告的 input tensor shape,只有 SDK 沒回報時才用這組值。原因是這裡的 + // 數字是人在上傳表單填的,實際案例是使用者填了 640x640 但模型根本不是 + // 那個尺寸 —— 尺寸錯了 NPU 不會報錯,只會安靜地給出錯的推論結果。 + // + // 零值 = 未宣告,bridge 端會忽略並往下 fallback。 + InputWidth int + InputHeight int +} + +// InferenceOptions 是推論期可即時調整的解析設定。 +// +// 與 FlashOptions 的分工:FlashOptions 在「把 model 載進裝置」時一次性帶入; +// InferenceOptions 則是在**同一個已載入的 model 上**改變輸出的解讀方式, +// 不需要重燒(KL520 重燒要數十秒)。 +// +// 兩個欄位的零值語意刻意不同,因為要能表達「不動」與「清空」兩種意圖: +// +// TaskType == "" → 不改變當前解析方式 +// Labels == nil → 不改變當前 label 表 +// Labels == []string{} → 清空 label 表,回到原始 enum(class_N) +// +// ⚠️ 因此 Labels 的判斷必須用 `!= nil` 而非 `len() > 0` —— 用長度判斷會讓 +// 「清空」這個合法意圖永遠送不出去。 +type InferenceOptions struct { + TaskType string + Labels []string +} + type FlashProgress struct { Percent int `json:"percent"` Stage string `json:"stage"` @@ -68,6 +115,9 @@ type InferenceResult struct { type ClassResult struct { Label string `json:"label"` Confidence float64 `json:"confidence"` + // ClassIndex 是模型輸出的原始類別索引,供前端在 label 缺漏時 fallback 顯示。 + // 刻意不加 omitempty —— index 0 是合法類別,omitempty 會把它吃掉。 + ClassIndex int `json:"classIndex"` } type DetectionResult struct { diff --git a/local-tool/server/internal/driver/kneron/kl720_driver.go b/local-tool/server/internal/driver/kneron/kl720_driver.go index 7f1399e..b29212c 100644 --- a/local-tool/server/internal/driver/kneron/kl720_driver.go +++ b/local-tool/server/internal/driver/kneron/kl720_driver.go @@ -508,6 +508,37 @@ func (d *KneronDriver) restartBridge() error { return nil } +// buildLoadModelCommand 組出 load_model 的 JSON-RPC payload。 +// +// ⚠️ Flash 有四處 load_model 呼叫點(初次 + 三條 retry 路徑)。四處都必須走這個 +// helper —— 如果任一處自己手寫 map,retry 成功後 task_type / labels 會遺失, +// 而且不會報錯(bridge 會 fallback 到檔名猜測,classification model 被誤判成 +// YOLO 只會回空結果)。這種失敗完全靜默,所以刻意集中在單一建構點。 +// +// 空值欄位不放進 payload:bridge 端把「缺欄位」與「空值」都當成未指定, +// 但少送欄位可讓 bridge log 的 "(not specified)" 語意精確。 +func buildLoadModelCommand(modelPath string, opts driver.FlashOptions) map[string]interface{} { + cmd := map[string]interface{}{ + "cmd": "load_model", + "path": modelPath, + } + if opts.TaskType != "" { + cmd["task_type"] = opts.TaskType + } + if len(opts.Labels) > 0 { + cmd["labels"] = opts.Labels + } + // 兩軸都要有值才送:只有一軸的宣告無法描述一個輸入尺寸,送過去只會讓 + // bridge 端多做一次驗證再丟掉。 + if opts.InputWidth > 0 && opts.InputHeight > 0 { + cmd["input_size"] = map[string]interface{}{ + "width": opts.InputWidth, + "height": opts.InputHeight, + } + } + return cmd +} + // Flash loads a model onto the Kneron device. Progress is reported through // the provided channel. // @@ -516,7 +547,10 @@ func (d *KneronDriver) restartBridge() error { // a full device reset + bridge restart + firmware reload. // - KL720 (flash-based): models can be freely reloaded. Error 40 // should not occur; if it does, a simple retry is attempted first. -func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashProgress) error { +// +// opts 帶著 models.json 宣告的 model metadata(taskType / labels),會隨每一次 +// load_model 送給 Python bridge —— 包含所有 retry 路徑。 +func (d *KneronDriver) Flash(modelPath string, opts driver.FlashOptions, progressCh chan<- driver.FlashProgress) error { d.mu.Lock() d.info.Status = driver.StatusFlashing pythonReady := d.pythonReady @@ -554,10 +588,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro } d.mu.Lock() - _, err := d.sendCommand(map[string]interface{}{ - "cmd": "load_model", - "path": modelPath, - }) + _, err := d.sendCommand(buildLoadModelCommand(modelPath, opts)) d.mu.Unlock() // Handle retryable errors (error 40, broken pipe). @@ -582,10 +613,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro } d.mu.Lock() - _, err = d.sendCommand(map[string]interface{}{ - "cmd": "load_model", - "path": modelPath, - }) + _, err = d.sendCommand(buildLoadModelCommand(modelPath, opts)) d.mu.Unlock() // If still failing, fall back to bridge restart as last resort. @@ -599,10 +627,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro } d.mu.Lock() d.info.Status = driver.StatusFlashing - _, err = d.sendCommand(map[string]interface{}{ - "cmd": "load_model", - "path": modelPath, - }) + _, err = d.sendCommand(buildLoadModelCommand(modelPath, opts)) d.mu.Unlock() } } else { @@ -626,10 +651,7 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro d.driverLog("INFO", "[kneron] bridge restarted, retrying load_model...") d.mu.Lock() d.info.Status = driver.StatusFlashing - _, err = d.sendCommand(map[string]interface{}{ - "cmd": "load_model", - "path": modelPath, - }) + _, err = d.sendCommand(buildLoadModelCommand(modelPath, opts)) d.mu.Unlock() } } @@ -697,6 +719,58 @@ func (d *KneronDriver) Flash(modelPath string, progressCh chan<- driver.FlashPro return nil } +// buildSetInferenceOptionsCommand 組出 set_inference_options 的 JSON-RPC payload。 +// +// 與 buildLoadModelCommand 的關鍵差異:這裡用「欄位在不在」表達意圖,所以 +// **不能**沿用「空值就不放進 payload」的規則 —— +// +// opts.TaskType == "" → 不帶 task_type 欄位 → bridge 保留當前解析方式 +// opts.Labels == nil → 不帶 labels 欄位 → bridge 保留當前 label 表 +// opts.Labels == [] → 帶空陣列 → bridge 清掉 label 表 +// +// 最後那條是刻意要能表達的狀態(使用者上傳錯 label 想清掉)。若照 load_model +// 的規則用 len()>0 判斷,「清掉」就永遠送不出去、變成靜默無效的操作。 +func buildSetInferenceOptionsCommand(opts driver.InferenceOptions) map[string]interface{} { + cmd := map[string]interface{}{ + "cmd": "set_inference_options", + } + if opts.TaskType != "" { + cmd["task_type"] = opts.TaskType + } + if opts.Labels != nil { + cmd["labels"] = opts.Labels + } + return cmd +} + +// SetInferenceOptions 在**不重新載入模型**的前提下,更新解析方式與 label 表。 +// +// KL520 一次只能載一個 model、換 model 要重燒(數十秒);但「怎麼解析輸出」 +// 與「index 顯示成什麼名字」都只是 post-process,可以即時切換。 +// +// 這個 method 刻意不放進 driver.DeviceDriver 介面 —— 它是 Kneron 特有能力, +// 放進去會逼三個既有 test fake 都跟著改。呼叫端改用窄介面 type-assert +// (同 firmware.UpgradeDriver 的做法)。 +func (d *KneronDriver) SetInferenceOptions(opts driver.InferenceOptions) error { + d.mu.Lock() + defer d.mu.Unlock() + + if !d.pythonReady { + return fmt.Errorf("hardware bridge is not running — device may not be connected") + } + if d.modelLoaded == "" { + return fmt.Errorf("no model loaded on device — flash a model first") + } + + if _, err := d.sendCommand(buildSetInferenceOptionsCommand(opts)); err != nil { + return fmt.Errorf("set inference options failed: %w", err) + } + + d.driverLog("INFO", "[kneron] inference options updated (taskType=%q, labels=%d)", + opts.TaskType, len(opts.Labels)) + return nil +} + // StartInference begins continuous inference mode. func (d *KneronDriver) StartInference() error { d.mu.Lock() diff --git a/local-tool/server/internal/driver/kneron/load_model_payload_test.go b/local-tool/server/internal/driver/kneron/load_model_payload_test.go new file mode 100644 index 0000000..3bde7ae --- /dev/null +++ b/local-tool/server/internal/driver/kneron/load_model_payload_test.go @@ -0,0 +1,313 @@ +package kneron + +import ( + "encoding/json" + "go/ast" + "go/parser" + "go/token" + "reflect" + "testing" + + "visiona-local/server/internal/driver" +) + +// TestBuildLoadModelCommand_WithMetadata:有帶 metadata 時 payload 含 +// task_type + labels,欄位名必須與 kneron_bridge.py handle_load_model 的 +// params.get("task_type") / params.get("labels") 完全一致。 +func TestBuildLoadModelCommand_WithMetadata(t *testing.T) { + cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{ + TaskType: "classification", + Labels: []string{"剪刀", "石頭", "布"}, + }) + + if cmd["cmd"] != "load_model" { + t.Errorf("cmd = %v, want load_model", cmd["cmd"]) + } + if cmd["path"] != "/models/rps.nef" { + t.Errorf("path = %v, want /models/rps.nef", cmd["path"]) + } + if cmd["task_type"] != "classification" { + t.Errorf("task_type = %v, want classification", cmd["task_type"]) + } + labels, ok := cmd["labels"].([]string) + if !ok { + t.Fatalf("labels type = %T, want []string", cmd["labels"]) + } + if !reflect.DeepEqual(labels, []string{"剪刀", "石頭", "布"}) { + t.Errorf("labels = %v, want [剪刀 石頭 布]", labels) + } +} + +// TestBuildLoadModelCommand_EmptyOptionsOmitsFields:沒帶 metadata 時不送 +// task_type / labels,讓 bridge 走既有 heuristics(既有 detection 行為不變)。 +func TestBuildLoadModelCommand_EmptyOptionsOmitsFields(t *testing.T) { + cmd := buildLoadModelCommand("/models/fcos.nef", driver.FlashOptions{}) + + if _, exists := cmd["task_type"]; exists { + t.Errorf("task_type should be omitted when empty, got %v", cmd["task_type"]) + } + if _, exists := cmd["labels"]; exists { + t.Errorf("labels should be omitted when empty, got %v", cmd["labels"]) + } + if len(cmd) != 2 { + t.Errorf("payload keys = %d (%v), want only cmd+path", len(cmd), cmd) + } +} + +// TestBuildLoadModelCommand_EmptyLabelsOmitted:labels 為 non-nil 空 slice 時 +// 也要省略 —— 送 [] 會讓 bridge 端 _sanitize_labels 走「空 list 視為未提供」, +// 語意雖同但多送無意義欄位。 +func TestBuildLoadModelCommand_EmptyLabelsOmitted(t *testing.T) { + cmd := buildLoadModelCommand("/m.nef", driver.FlashOptions{ + TaskType: "object_detection", + Labels: []string{}, + }) + if _, exists := cmd["labels"]; exists { + t.Errorf("empty labels should be omitted, got %v", cmd["labels"]) + } + if cmd["task_type"] != "object_detection" { + t.Errorf("task_type = %v, want object_detection", cmd["task_type"]) + } +} + +// TestBuildLoadModelCommand_SerializesToBridgeContract:payload 經 JSON 編碼後 +// 必須是 bridge 能吃的形狀(labels 是 JSON array of string,不是物件)。 +// sendCommand 實際就是把 map 丟給 json.Marshal 送進 stdin。 +func TestBuildLoadModelCommand_SerializesToBridgeContract(t *testing.T) { + cmd := buildLoadModelCommand("/m.nef", driver.FlashOptions{ + TaskType: "classification", + Labels: []string{"a", "b"}, + }) + + data, err := json.Marshal(cmd) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } + + var decoded map[string]interface{} + if err := json.Unmarshal(data, &decoded); err != nil { + t.Fatalf("unmarshal failed: %v", err) + } + if decoded["task_type"] != "classification" { + t.Errorf("task_type after roundtrip = %v", decoded["task_type"]) + } + rawLabels, ok := decoded["labels"].([]interface{}) + if !ok { + t.Fatalf("labels after roundtrip type = %T, want JSON array", decoded["labels"]) + } + if len(rawLabels) != 2 || rawLabels[0] != "a" || rawLabels[1] != "b" { + t.Errorf("labels after roundtrip = %v, want [a b]", rawLabels) + } +} + +// TestFlashUsesBuilderForEveryLoadModelCall 是本次改動最重要的一條測試。 +// +// Flash 有四處 load_model 呼叫點(初次 + KL720 簡單 retry + KL720 restart retry +// + KL520 restart retry)。漏改任一處的後果是「retry 成功後 model metadata 靜默 +// 遺失」—— 不會報錯、不會 panic,只會讓 classification model 被誤判成 YOLO 而 +// 回傳空結果,極難從現象追回根因(plan §7 R-5)。 +// +// 用 AST 掃 Flash 函式本體,斷言: +// 1. 函式內沒有任何自己手寫的 load_model map literal +// 2. 所有 sendCommand 的 load_model 都經過 buildLoadModelCommand +// 3. 呼叫點數量 == 4(將來新增 retry 路徑忘了帶 opts 時,這條會亮) +func TestFlashUsesBuilderForEveryLoadModelCall(t *testing.T) { + fset := token.NewFileSet() + file, err := parser.ParseFile(fset, "kl720_driver.go", nil, 0) + if err != nil { + t.Fatalf("parse kl720_driver.go: %v", err) + } + + var flashFn *ast.FuncDecl + for _, decl := range file.Decls { + fn, ok := decl.(*ast.FuncDecl) + if !ok || fn.Name.Name != "Flash" || fn.Recv == nil { + continue + } + flashFn = fn + break + } + if flashFn == nil { + t.Fatal("Flash method not found in kl720_driver.go") + } + + builderCalls := 0 + rawLiterals := 0 + + ast.Inspect(flashFn, func(n ast.Node) bool { + switch node := n.(type) { + case *ast.CallExpr: + if ident, ok := node.Fun.(*ast.Ident); ok && ident.Name == "buildLoadModelCommand" { + builderCalls++ + } + case *ast.CompositeLit: + // 偵測 Flash 內自己手寫的 map,其中含 "load_model" 字串。 + for _, elt := range node.Elts { + kv, ok := elt.(*ast.KeyValueExpr) + if !ok { + continue + } + lit, ok := kv.Value.(*ast.BasicLit) + if ok && lit.Kind == token.STRING && lit.Value == `"load_model"` { + rawLiterals++ + } + } + } + return true + }) + + if rawLiterals != 0 { + t.Errorf("Flash 內有 %d 個手寫的 load_model map literal;"+ + "所有呼叫點都必須走 buildLoadModelCommand,否則 retry 後 "+ + "task_type/labels 會靜默遺失", rawLiterals) + } + + const wantCallSites = 4 + if builderCalls != wantCallSites { + t.Errorf("buildLoadModelCommand 呼叫點 = %d,want %d "+ + "(初次 + KL720 retry + KL720 restart retry + KL520 restart retry)。"+ + "若確實新增/移除了 retry 路徑,請確認新路徑有帶 opts 後再更新此數字", + builderCalls, wantCallSites) + } +} + +// TestParseInferenceResult_ClassIndexPreserved:bridge 回傳的 classIndex 要 +// 進得了 Go struct。加 ClassIndex 欄位前,這個值會被 encoding/json 靜默丟棄。 +func TestParseInferenceResult_ClassIndexPreserved(t *testing.T) { + resp := map[string]interface{}{ + "taskType": "classification", + "timestamp": float64(1721545200000), + "latencyMs": 45.2, + "classifications": []interface{}{ + map[string]interface{}{"label": "石頭", "confidence": 0.94, "classIndex": float64(1)}, + map[string]interface{}{"label": "class_0", "confidence": 0.04, "classIndex": float64(0)}, + }, + } + + result, err := parseInferenceResult(resp) + if err != nil { + t.Fatalf("parseInferenceResult failed: %v", err) + } + if result.TaskType != "classification" { + t.Errorf("TaskType = %q, want classification", result.TaskType) + } + if len(result.Classifications) != 2 { + t.Fatalf("Classifications = %d, want 2", len(result.Classifications)) + } + if result.Classifications[0].ClassIndex != 1 { + t.Errorf("Classifications[0].ClassIndex = %d, want 1", result.Classifications[0].ClassIndex) + } + // index 0 是合法類別 —— 若 struct tag 誤加 omitempty,這筆會在序列化時消失。 + if result.Classifications[1].ClassIndex != 0 { + t.Errorf("Classifications[1].ClassIndex = %d, want 0", result.Classifications[1].ClassIndex) + } +} + +// TestClassResult_ClassIndexZeroNotOmitted:index 0 必須出現在送給前端的 JSON。 +// 這是 omitempty 會踩的陷阱 —— 前端拿不到 classIndex 就無法做 fallback 顯示。 +func TestClassResult_ClassIndexZeroNotOmitted(t *testing.T) { + data, err := json.Marshal(driver.ClassResult{Label: "class_0", Confidence: 0.9, ClassIndex: 0}) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } + + var decoded map[string]interface{} + if err := json.Unmarshal(data, &decoded); err != nil { + t.Fatalf("unmarshal failed: %v", err) + } + if _, exists := decoded["classIndex"]; !exists { + t.Errorf("classIndex 不見了(大概是加了 omitempty):%s", data) + } + if decoded["classIndex"] != float64(0) { + t.Errorf("classIndex = %v, want 0", decoded["classIndex"]) + } +} + +// TestBuildLoadModelCommand_InputSizeIncluded:宣告的 input size 要送到 bridge, +// 欄位名與巢狀結構必須與 kneron_bridge.py 的 +// _normalize_declared_input_size(params.get("input_size")) 一致。 +func TestBuildLoadModelCommand_InputSizeIncluded(t *testing.T) { + cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{ + TaskType: "classification", + InputWidth: 320, + InputHeight: 256, + }) + + size, ok := cmd["input_size"].(map[string]interface{}) + if !ok { + t.Fatalf("input_size type = %T, want map[string]interface{}", cmd["input_size"]) + } + if size["width"] != 320 { + t.Errorf("input_size.width = %v, want 320", size["width"]) + } + // 高度不可被壓成寬度 —— 非正方形模型兩軸必須各自送出。 + if size["height"] != 256 { + t.Errorf("input_size.height = %v, want 256", size["height"]) + } +} + +// TestBuildLoadModelCommand_ZeroInputSizeOmitted:models.json 沒填 inputSize 時 +// 兩軸都是 0,送 0 過去只會讓 bridge 多驗一次再丟掉,且會讓 log 的 +// "(not specified)" 語意失真。 +func TestBuildLoadModelCommand_ZeroInputSizeOmitted(t *testing.T) { + cmd := buildLoadModelCommand("/models/fcos.nef", driver.FlashOptions{ + TaskType: "object_detection", + }) + + if _, exists := cmd["input_size"]; exists { + t.Errorf("input_size should be omitted when zero, got %v", cmd["input_size"]) + } +} + +// TestBuildLoadModelCommand_PartialInputSizeOmitted:只有一軸的宣告無法描述 +// 一個輸入尺寸。半套送出比不送更危險 —— bridge 端會看到一個看似有效的來源。 +func TestBuildLoadModelCommand_PartialInputSizeOmitted(t *testing.T) { + for _, tc := range []struct { + name string + width int + height int + }{ + {"only width", 320, 0}, + {"only height", 0, 320}, + {"negative width", -1, 320}, + } { + t.Run(tc.name, func(t *testing.T) { + cmd := buildLoadModelCommand("/models/m.nef", driver.FlashOptions{ + InputWidth: tc.width, + InputHeight: tc.height, + }) + if _, exists := cmd["input_size"]; exists { + t.Errorf("input_size should be omitted, got %v", cmd["input_size"]) + } + }) + } +} + +// TestBuildLoadModelCommand_InputSizeSerializesToBridgeShape:payload 實際被 +// JSON 序列化後的形狀,就是 bridge 端會 parse 到的東西。用序列化後的結果斷言 +// 可避免「Go 端看起來對、上 wire 後欄位名或巢狀層級不同」。 +func TestBuildLoadModelCommand_InputSizeSerializesToBridgeShape(t *testing.T) { + cmd := buildLoadModelCommand("/models/rps.nef", driver.FlashOptions{ + InputWidth: 320, + InputHeight: 320, + }) + + data, err := json.Marshal(cmd) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } + + var decoded struct { + InputSize struct { + Width int `json:"width"` + Height int `json:"height"` + } `json:"input_size"` + } + if err := json.Unmarshal(data, &decoded); err != nil { + t.Fatalf("unmarshal failed: %v", err) + } + if decoded.InputSize.Width != 320 || decoded.InputSize.Height != 320 { + t.Errorf("input_size = %dx%d, want 320x320 (payload: %s)", + decoded.InputSize.Width, decoded.InputSize.Height, data) + } +} diff --git a/local-tool/server/internal/flash/service.go b/local-tool/server/internal/flash/service.go index 676cce2..ad35651 100644 --- a/local-tool/server/internal/flash/service.go +++ b/local-tool/server/internal/flash/service.go @@ -12,6 +12,28 @@ import ( "visiona-local/server/internal/model" ) +// 可指定的推論種類。與 models.json / 前端同一組值。 +// +// models.json 另有 segmentation / pose_estimation,但 Python bridge 與前端都還 +// 沒有對應的解析路徑,所以只開放這兩種真的能產出結果的種類。 +// +// 燒錄時不再讓使用者選推論種類(改由推論期的 +// POST /devices/:id/inference/options 即時切換、不必重燒),但這組常數與 +// IsValidTaskTypeOverride 仍是「解析方式」的值域來源,由該 endpoint 沿用, +// 讓 wire 上永遠只有一組合法命名。 +const ( + TaskTypeClassification = "classification" + TaskTypeObjectDetection = "object_detection" +) + +// IsValidTaskTypeOverride 回報 taskType 是否為合法的解析方式覆寫值。 +// +// 空字串(未指定)不算合法覆寫 —— 呼叫端要自己先判斷「有沒有要覆寫」, +// 這樣「未指定」與「指定了但打錯字」不會被混為一談。 +func IsValidTaskTypeOverride(taskType string) bool { + return taskType == TaskTypeClassification || taskType == TaskTypeObjectDetection +} + func isCompatible(modelHardware []string, deviceType string) bool { dt := strings.ToUpper(deviceType) for _, hw := range modelHardware { @@ -84,6 +106,11 @@ func (s *Service) CleanupTask(taskID string) { s.tracker.Remove(taskID) } +// StartFlash 把 model 載入到裝置。 +// +// 推論種類一律用 models.json 宣告的值。使用者若要改解析方式,走推論期的 +// POST /devices/:id/inference/options —— 那條路徑不必重燒、可即時切換, +// 功能完全涵蓋燒錄時再選一次的舊做法。 func (s *Service) StartFlash(deviceID, modelID string) (string, <-chan driver.FlashProgress, error) { session, err := s.deviceMgr.GetDevice(deviceID) if err != nil { @@ -137,7 +164,18 @@ func (s *Service) StartFlash(deviceID, modelID string) (string, <-chan driver.Fl time.Sleep(500 * time.Millisecond) - flashErr := session.Driver.Flash(modelPath, task.ProgressCh) + // 把 models.json 宣告的 metadata 一起帶下去 —— bridge 端有 taskType + // 就不再靠檔名猜 model type(自訂模型存成 model.nef、檔名沒有關鍵字, + // 猜測必定落到 detection 分支)。labels 純顯示層、沒有也能跑。 + // + // inputSize 是宣告值、**優先序最低**:bridge 端會先問 SDK 模型自己 + // 宣告的 input shape,只有問不到才用這裡的值(這欄是人填的,可能亂填)。 + flashErr := session.Driver.Flash(modelPath, driver.FlashOptions{ + TaskType: m.TaskType, + Labels: m.Labels, + InputWidth: m.InputSize.Width, + InputHeight: m.InputSize.Height, + }, task.ProgressCh) // Flash 完成或失敗後,driver 不會再寫 progressCh,安全地寫 error 訊息然後 close。 if flashErr != nil { diff --git a/local-tool/server/internal/flash/service_metadata_test.go b/local-tool/server/internal/flash/service_metadata_test.go new file mode 100644 index 0000000..61b8ff6 --- /dev/null +++ b/local-tool/server/internal/flash/service_metadata_test.go @@ -0,0 +1,190 @@ +package flash + +import ( + "bytes" + "go/ast" + "go/parser" + "go/printer" + "go/token" + "reflect" + "strings" + "testing" + + "visiona-local/server/internal/driver" + "visiona-local/server/internal/model" +) + +// recordingDriver 記錄 Flash 收到的 opts,用來驗證 model metadata 有被傳下去。 +type recordingDriver struct { + gotPath string + gotOpts driver.FlashOptions + called bool +} + +func (d *recordingDriver) Info() driver.DeviceInfo { return driver.DeviceInfo{} } +func (d *recordingDriver) Connect() error { return nil } +func (d *recordingDriver) Disconnect() error { return nil } +func (d *recordingDriver) IsConnected() bool { return true } +func (d *recordingDriver) Flash(modelPath string, opts driver.FlashOptions, _ chan<- driver.FlashProgress) error { + d.called = true + d.gotPath = modelPath + d.gotOpts = opts + return nil +} +func (d *recordingDriver) StartInference() error { return nil } +func (d *recordingDriver) StopInference() error { return nil } +func (d *recordingDriver) ReadInference() (*driver.InferenceResult, error) { return nil, nil } +func (d *recordingDriver) RunInference(_ []byte) (*driver.InferenceResult, error) { + return nil, nil +} +func (d *recordingDriver) GetModelInfo() (*driver.ModelInfo, error) { return nil, nil } + +// flashOptionsFor 複製 StartFlash 內部組 FlashOptions 的邏輯。 +// +// 為什麼不直接跑 StartFlash:Service 依賴具體的 *device.Manager,而 Manager 的 +// sessions map 未匯出、只能由真實硬體偵測填入,沒有注入 fake session 的接縫。 +// 為了測試而改 production 的依賴結構超出 M2 範圍,所以這裡改為釘住「送進 +// driver 的 FlashOptions 必須完整帶著 model 的 TaskType/Labels」這個契約。 +func flashOptionsFor(m model.Model) driver.FlashOptions { + return driver.FlashOptions{ + TaskType: m.TaskType, + Labels: m.Labels, + InputWidth: m.InputSize.Width, + InputHeight: m.InputSize.Height, + } +} + +// TestFlashOptions_CarriesClassificationMetadata:classification model 的 +// taskType + labels 要完整傳到 driver,不能像改動前一樣只傳 path 就丟棄。 +func TestFlashOptions_CarriesClassificationMetadata(t *testing.T) { + m := model.Model{ + ID: "custom-rps", + TaskType: "classification", + Labels: []string{"剪刀", "石頭", "布"}, + } + + d := &recordingDriver{} + opts := flashOptionsFor(m) + if err := d.Flash("/models/custom-rps/model.nef", opts, nil); err != nil { + t.Fatalf("Flash returned error: %v", err) + } + + if !d.called { + t.Fatal("Flash was not called") + } + if d.gotOpts.TaskType != "classification" { + t.Errorf("TaskType = %q, want classification", d.gotOpts.TaskType) + } + if !reflect.DeepEqual(d.gotOpts.Labels, []string{"剪刀", "石頭", "布"}) { + t.Errorf("Labels = %v, want [剪刀 石頭 布]", d.gotOpts.Labels) + } +} + +// TestFlashOptions_CarriesDetectionMetadata:既有 detection model 走同一條路, +// taskType 為 object_detection —— bridge 端收到後仍走 detection 分支,行為不變。 +func TestFlashOptions_CarriesDetectionMetadata(t *testing.T) { + m := model.Model{ + ID: "kl520-fcos-detection", + TaskType: "object_detection", + Labels: []string{"person", "bicycle", "car"}, + } + + opts := flashOptionsFor(m) + if opts.TaskType != "object_detection" { + t.Errorf("TaskType = %q, want object_detection", opts.TaskType) + } + if len(opts.Labels) != 3 { + t.Errorf("Labels = %v, want 3 entries", opts.Labels) + } +} + +// TestFlashOptions_EmptyMetadataIsZeroValue:model 沒宣告 taskType/labels 時 +// 送出的是零值,driver 端會據此省略欄位,讓 bridge fallback 到既有 heuristics。 +func TestFlashOptions_EmptyMetadataIsZeroValue(t *testing.T) { + opts := flashOptionsFor(model.Model{ID: "bare"}) + + if opts.TaskType != "" { + t.Errorf("TaskType = %q, want empty", opts.TaskType) + } + if len(opts.Labels) != 0 { + t.Errorf("Labels = %v, want empty", opts.Labels) + } +} + +// TestStartFlashPassesModelMetadata 釘住 StartFlash 原始碼真的有把 m.TaskType / +// m.Labels 傳進 Flash。 +// +// 上面的測試只驗「FlashOptions 帶得動 metadata」,無法防止有人把 service.go 改回 +// 只傳 path —— 那正是改動前的 bug 形態(metadata 被靜默丟棄、不會報錯)。 +// 這裡直接掃 service.go 的 StartFlash 本體補上這個缺口。 +func TestStartFlashPassesModelMetadata(t *testing.T) { + src := readStartFlashSource(t) + + for _, want := range []string{ + "m.TaskType", "m.Labels", "driver.FlashOptions", + // 宣告的 input size 也要傳下去 —— 沒傳的話 bridge 在 SDK 問不到 + // shape 時只能靠檔名猜,那正是「尺寸靜默錯誤」的來源。 + "m.InputSize.Width", "m.InputSize.Height", + } { + if !strings.Contains(src, want) { + t.Errorf("StartFlash 原始碼缺少 %q —— model metadata 沒有被傳給 driver.Flash", want) + } + } +} + +// TestFlashOptions_CarriesDeclaredInputSize:models.json 宣告的 inputSize 要 +// 完整帶到 driver。改動前這個欄位在推論鏈路上完全沒被讀過 —— 使用者在上傳 +// 表單填的寬高毫無作用,input size 全由檔名猜測決定。 +func TestFlashOptions_CarriesDeclaredInputSize(t *testing.T) { + m := model.Model{ + ID: "custom-rps", + TaskType: "classification", + InputSize: model.InputSize{Width: 320, Height: 256}, + } + + opts := flashOptionsFor(m) + if opts.InputWidth != 320 { + t.Errorf("InputWidth = %d, want 320", opts.InputWidth) + } + // 非正方形模型的高度不可被壓成寬度。 + if opts.InputHeight != 256 { + t.Errorf("InputHeight = %d, want 256", opts.InputHeight) + } +} + +// TestFlashOptions_MissingInputSizeIsZero:models.json 沒填 inputSize 時是零值, +// driver 端會省略欄位,bridge 據此 fallback(既有 detection 模型走這條)。 +func TestFlashOptions_MissingInputSizeIsZero(t *testing.T) { + opts := flashOptionsFor(model.Model{ID: "bare"}) + + if opts.InputWidth != 0 || opts.InputHeight != 0 { + t.Errorf("InputSize = %dx%d, want 0x0", + opts.InputWidth, opts.InputHeight) + } +} + +// readStartFlashSource 取出 service.go 中 StartFlash 方法的原始碼文字。 +func readStartFlashSource(t *testing.T) string { + t.Helper() + + fset := token.NewFileSet() + file, err := parser.ParseFile(fset, "service.go", nil, 0) + if err != nil { + t.Fatalf("parse service.go: %v", err) + } + + for _, decl := range file.Decls { + fn, ok := decl.(*ast.FuncDecl) + if !ok || fn.Name.Name != "StartFlash" || fn.Recv == nil { + continue + } + var buf bytes.Buffer + if err := printer.Fprint(&buf, fset, fn); err != nil { + t.Fatalf("print StartFlash: %v", err) + } + return buf.String() + } + + t.Fatal("StartFlash method not found in service.go") + return "" +}