visionA/local-tool/server/internal/driver/kneron/load_model_payload_test.go
jim800121chen ddd1aae5d1 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>
2026-07-22 19:27:02 +08:00

314 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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