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>
314 lines
11 KiB
Go
314 lines
11 KiB
Go
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)
|
||
}
|
||
}
|