220 lines
6.6 KiB
Go
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")
|
|
}
|
|
}
|