81 lines
2.7 KiB
Go
81 lines
2.7 KiB
Go
package ai
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
type currentASRWireCapture struct {
|
|
path, key string
|
|
frame []byte
|
|
err error
|
|
}
|
|
|
|
func TestCurrentASRParametersReachDoubaoSDKWebSocket(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
mode, resultType string
|
|
nonstream bool
|
|
}{
|
|
{"full_ai", "full", false},
|
|
{"asr_only", "single", true},
|
|
} {
|
|
t.Run(tc.mode, func(t *testing.T) {
|
|
captures := make(chan currentASRWireCapture, 1)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
conn, err := (&websocket.Upgrader{}).Upgrade(w, r, nil)
|
|
if err != nil {
|
|
captures <- currentASRWireCapture{err: err}
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
_, frame, err := conn.ReadMessage()
|
|
captures <- currentASRWireCapture{path: r.URL.Path, key: r.Header.Get("X-Api-Key"), frame: frame, err: err}
|
|
}))
|
|
defer server.Close()
|
|
task, providers := currentFixture(t, tc.mode)
|
|
asrProvider := providers["asr-example"]
|
|
asrProvider.Endpoint = server.URL
|
|
providers[asrProvider.ProviderRef] = asrProvider
|
|
bound, err := BindCurrent(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
if _, err := bound.Recognize(ctx, []byte{1, 0}); err == nil {
|
|
t.Fatal("disconnected mock ASR cannot pretend recognition completed")
|
|
}
|
|
select {
|
|
case got := <-captures:
|
|
if got.err != nil || got.path != "/api/v3/sauc/bigmodel" || got.key != asrProvider.Credential {
|
|
t.Fatalf("SDK ASR WebSocket endpoint/credential was changed: path=%q err=%v", got.path, got.err)
|
|
}
|
|
// Inspect the SDK-emitted JSON payload. No provider protocol is
|
|
// implemented by the application or by this mock responder.
|
|
start := bytes.IndexByte(got.frame, '{')
|
|
if start < 0 {
|
|
t.Fatal("SDK sent no ASR start JSON frame")
|
|
}
|
|
var payload map[string]any
|
|
if err := json.Unmarshal(got.frame[start:], &payload); err != nil {
|
|
t.Fatalf("SDK ASR start JSON cannot be decoded: %v", err)
|
|
}
|
|
audio := payload["audio"].(map[string]any)
|
|
request := payload["request"].(map[string]any)
|
|
if audio["format"] != "pcm_s16le" || audio["sample_rate"] != float64(16000) || audio["channel"] != float64(1) || audio["bits"] != float64(16) || audio["language"] != "zh-CN" || request["result_type"] != tc.resultType || request["enable_nonstream"] != tc.nonstream {
|
|
t.Fatalf("SDK ASR request omitted approved input/interim fields: audio=%v request=%v", audio, request)
|
|
}
|
|
case <-ctx.Done():
|
|
t.Fatal("SDK did not send an ASR start request to isolated mock")
|
|
}
|
|
})
|
|
}
|
|
}
|