package main import ( "bytes" "compress/gzip" "context" "encoding/binary" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/gorilla/websocket" ) const testToken = "0123456789abcdef0123456789abcdef" func TestModelCredentialRouting(t *testing.T) { t.Setenv("BAILIAN_ASR_MODELS", "fun-asr-realtime,qwen-audio-asr-realtime,volc-bigmodel") s := &server{cfg: config{funASRKey: "test-fun-key", funASRBase: "wss://fun.test", bailianBase: "wss://bailian.test"}} models := s.models() if len(models) != 3 || !models[0].Enabled || models[1].Enabled || models[2].Enabled { t.Fatalf("credential availability: %+v", models) } if _, err := s.provider("qwen-audio-asr-realtime"); err == nil { t.Fatal("missing Bailian credentials accepted") } key, endpoint := s.bailianCredentials("fun-asr-realtime") if key != "test-fun-key" || endpoint != "wss://fun.test" { t.Fatal("Fun-ASR routing ignored explicit configuration") } } func TestHTTPAuthenticationAndCapabilities(t *testing.T) { s := &server{token: testToken, slots: make(chan struct{}, 1)} h := s.handler(context.Background()) for _, path := range []string{"/", "/healthz", "/api/models"} { r := httptest.NewRequest("GET", path, nil) w := httptest.NewRecorder() h.ServeHTTP(w, r) want := 200 if path == "/api/models" { want = 401 } if w.Code != want { t.Fatalf("%s: %d", path, w.Code) } } r := httptest.NewRequest("GET", "/api/models", nil) r.Header.Set("Authorization", "Bearer "+testToken) w := httptest.NewRecorder() h.ServeHTTP(w, r) if w.Code != 200 { t.Fatal(w.Code) } r.Header.Set("Origin", "https://evil.invalid") if s.sameOrigin(r) { t.Fatal("cross origin accepted") } if err := validateConfig("short", "", loadConfig()); err == nil { t.Fatal("short token accepted") } if err := validateConfig(testToken, "https://host.invalid/path", loadConfig()); err == nil { t.Fatal("invalid origin accepted") } } func TestWebASRPipelineAndDisconnect(t *testing.T) { upstreamClosed := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { up := websocket.Upgrader{} c, err := up.Upgrade(w, r, nil) if err != nil { return } defer c.Close() defer close(upstreamClosed) var req map[string]any if c.ReadJSON(&req) != nil { return } if c.WriteJSON(map[string]any{"header": map[string]string{"event": "task-started"}}) != nil { return } for { kind, _, err := c.ReadMessage() if err != nil { return } if kind == websocket.BinaryMessage { _ = c.WriteJSON(map[string]any{"header": map[string]string{"event": "result-generated"}, "payload": map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "离线模拟识别", "sentence_end": true}}}}) } } })) defer upstream.Close() s := &server{token: testToken, slots: make(chan struct{}, 1), newProvider: func(id string) (asrProvider, error) { return newBailianASR(id, agentCfg{BailianKey: "fake", BailianWssBaseURL: "ws" + strings.TrimPrefix(upstream.URL, "http")}), nil }} httpServer := httptest.NewServer(s.handler(context.Background())) defer httpServer.Close() dialer := websocket.Dialer{Subprotocols: []string{"asr.v1", "auth." + testToken}} url := "ws" + strings.TrimPrefix(httpServer.URL, "http") + "/ws" _, response, err := dialer.Dial(url, http.Header{"Origin": []string{"https://evil.invalid"}}) if err == nil || response.StatusCode != 403 { t.Fatal("cross-origin websocket not rejected") } c, _, err := dialer.Dial(url, nil) if err != nil { t.Fatal(err) } defer c.Close() _ = c.SetReadDeadline(time.Now().Add(3 * time.Second)) var event map[string]any if err = c.ReadJSON(&event); err != nil || event["type"] != "ready" { t.Fatalf("ready: %v %v", event, err) } if err = c.WriteJSON(map[string]string{"type": "start", "model": "test"}); err != nil { t.Fatal(err) } if err = c.ReadJSON(&event); err != nil || event["type"] != "asr-started" { t.Fatalf("start: %v %v", event, err) } if err = c.WriteMessage(websocket.BinaryMessage, make([]byte, 3200)); err != nil { t.Fatal(err) } if err = c.ReadJSON(&event); err != nil || event["type"] != "final" || event["text"] != "离线模拟识别" { t.Fatalf("final: %v %v", event, err) } _ = c.Close() select { case <-upstreamClosed: case <-time.After(2 * time.Second): t.Fatal("disconnect leaked upstream socket") } } func TestStartCancellation(t *testing.T) { ready := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { up := websocket.Upgrader{} c, err := up.Upgrade(w, r, nil) if err != nil { return } defer c.Close() _, _, _ = c.ReadMessage() close(ready) _, _, _ = c.ReadMessage() })) defer upstream.Close() a := newBailianASR("test", agentCfg{BailianWssBaseURL: "ws" + strings.TrimPrefix(upstream.URL, "http")}) ctx, cancel := context.WithCancel(context.Background()) defer cancel() result := make(chan error, 1) go func() { result <- a.start(ctx) }() <-ready cancel() select { case err := <-result: if err == nil { t.Fatal("canceled start succeeded") } case <-time.After(2 * time.Second): t.Fatal("start ignored cancellation") } } func TestFinalEventBackpressureAndClose(t *testing.T) { a := newBailianASR("test", agentCfg{}) for i := 0; i < cap(a.out); i++ { a.out <- asrEvent{Typ: "partial"} } done := make(chan struct{}) go func() { a.emit(asrEvent{Typ: "final", Text: "must keep"}); close(done) }() <-a.out select { case <-done: case <-time.After(time.Second): t.Fatal("final delivery blocked") } found := false for len(a.out) > 0 { if (<-a.out).Typ == "final" { found = true } } if !found { t.Fatal("final event silently dropped") } a.close() a.close() a.closeUpstream() } func TestVolcFrameBounds(t *testing.T) { frame := volcFrame(msgFullServer, flagLast, serJSON, compNone, []byte(`{"result":{}}`)) typ, flags, _, payload, err := decodeVolc(frame) if err != nil || typ != msgFullServer || flags != flagLast || string(payload) != `{"result":{}}` { t.Fatal("last frame without sequence rejected", err) } for n := 0; n < len(frame); n++ { if _, _, _, _, err := decodeVolc(frame[:n]); err == nil { t.Fatalf("truncated length %d accepted", n) } } bad := append([]byte(nil), frame...) binary.BigEndian.PutUint32(bad[4:8], 0xffffffff) if _, _, _, _, err := decodeVolc(bad); err == nil { t.Fatal("oversized length accepted") } var b bytes.Buffer w := gzip.NewWriter(&b) _, _ = w.Write(make([]byte, (1<<20)+1)) _ = w.Close() if _, _, _, _, err := decodeVolc(volcFrame(msgFullServer, 0, serJSON, compGzip, b.Bytes())); err == nil { t.Fatal("gzip expansion not bounded") } }