Files
2026-09-11 17:47:03 +08:00

220 lines
6.6 KiB
Go

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