feat(server): 把 model metadata 傳進推論鏈路

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) <noreply@anthropic.com>
This commit is contained in:
jim800121chen 2026-07-22 19:27:02 +08:00
parent f14d24bd7b
commit ddd1aae5d1
7 changed files with 690 additions and 21 deletions

View File

@ -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) {

View File

@ -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 }

View File

@ -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 輸出原始 enumclass_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 表,回到原始 enumclass_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 {

View File

@ -508,6 +508,37 @@ func (d *KneronDriver) restartBridge() error {
return nil
}
// buildLoadModelCommand 組出 load_model 的 JSON-RPC payload。
//
// ⚠️ Flash 有四處 load_model 呼叫點(初次 + 三條 retry 路徑)。四處都必須走這個
// helper —— 如果任一處自己手寫 mapretry 成功後 task_type / labels 會遺失,
// 而且不會報錯bridge 會 fallback 到檔名猜測classification model 被誤判成
// YOLO 只會回空結果)。這種失敗完全靜默,所以刻意集中在單一建構點。
//
// 空值欄位不放進 payloadbridge 端把「缺欄位」與「空值」都當成未指定,
// 但少送欄位可讓 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 metadatataskType / 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()

View File

@ -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_EmptyLabelsOmittedlabels 為 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_SerializesToBridgeContractpayload 經 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 呼叫點 = %dwant %d "+
"(初次 + KL720 retry + KL720 restart retry + KL520 restart retry)。"+
"若確實新增/移除了 retry 路徑,請確認新路徑有帶 opts 後再更新此數字",
builderCalls, wantCallSites)
}
}
// TestParseInferenceResult_ClassIndexPreservedbridge 回傳的 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_ClassIndexZeroNotOmittedindex 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_ZeroInputSizeOmittedmodels.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_InputSizeSerializesToBridgeShapepayload 實際被
// 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)
}
}

View File

@ -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 {

View File

@ -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 的邏輯。
//
// 為什麼不直接跑 StartFlashService 依賴具體的 *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_CarriesClassificationMetadataclassification 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_EmptyMetadataIsZeroValuemodel 沒宣告 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_CarriesDeclaredInputSizemodels.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_MissingInputSizeIsZeromodels.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 ""
}