Files
go-sip/internal/rpc/approved_runner_test.go
T

104 lines
3.4 KiB
Go

package rpc
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"os"
"strings"
"sync/atomic"
"testing"
"time"
"git.ipao.vip/rogee/go-sip/internal/ai"
"git.ipao.vip/rogee/go-sip/internal/configread"
"git.ipao.vip/rogee/go-sip/internal/media"
)
type blockedApprovedMedia struct {
sent int
rate int
}
func (*blockedApprovedMedia) ReadPayload(ctx context.Context) ([]byte, error) {
<-ctx.Done()
return nil, ctx.Err()
}
func (m *blockedApprovedMedia) SendPCM16(ctx context.Context, pcm []byte, rate int) error {
if err := ctx.Err(); err != nil {
return err
}
m.sent++
m.rate = rate
return nil
}
func (*blockedApprovedMedia) Stats() media.RTPStats { return media.RTPStats{} }
func TestApprovedCallRunnerASROnlyHonorsSignedTimeout(t *testing.T) {
req := approvedTestRequest(t, time.Date(2026, 9, 20, 10, 0, 0, 0, time.UTC))
var task configread.CurrentTask
if err := json.Unmarshal(req.TaskConfigJson, &task); err != nil {
t.Fatal(err)
}
var providers map[string]configread.CurrentProvider
if err := json.Unmarshal(req.ProvidersJson, &providers); err != nil {
t.Fatal(err)
}
bound, err := ai.BindCurrent(task, providers)
if err != nil {
t.Fatal(err)
}
mediaSession := &blockedApprovedMedia{}
_, err = RunApprovedCall(context.Background(), ApprovedExecution{AI: bound, MaxCallDuration: 35 * time.Millisecond}, mediaSession, func(context.Context) error { return nil })
if err == nil || !strings.Contains(err.Error(), "deadline") || mediaSession.sent != 0 {
t.Fatalf("ASR-only call must timeout without opening or outbound TTS: sent=%d err=%v", mediaSession.sent, err)
}
}
func TestApprovedCallRunnerSynthesizesOpeningBeforeMediaCapture(t *testing.T) {
raw, err := os.ReadFile("../../contracts/local/examples/config-read-task-full.json")
if err != nil {
t.Fatal(err)
}
var task configread.CurrentTask
if err := json.Unmarshal(raw, &task); err != nil {
t.Fatal(err)
}
req := approvedTestRequest(t, time.Date(2026, 9, 20, 10, 0, 0, 0, time.UTC))
var providers map[string]configread.CurrentProvider
if err := json.Unmarshal(req.ProvidersJson, &providers); err != nil {
t.Fatal(err)
}
var sdkCalls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sdkCalls.Add(1)
var payload struct {
Params struct {
Text string `json:"text"`
} `json:"req_params"`
}
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil || payload.Params.Text != "Example greeting" {
http.Error(w, "opening text changed", http.StatusBadRequest)
return
}
_, _ = fmt.Fprintf(w, `{"code":0,"data":%q}`+"\n", base64.StdEncoding.EncodeToString([]byte{1, 0, 2, 0}))
_, _ = fmt.Fprintln(w, `{"code":20000000,"message":"ok","data":null}`)
}))
defer server.Close()
tts := providers["tts-example"]
tts.Endpoint = server.URL
providers[tts.ProviderRef] = tts
bound, err := ai.BindCurrent(task, providers)
if err != nil {
t.Fatal(err)
}
mediaSession := &blockedApprovedMedia{}
_, err = RunApprovedCall(context.Background(), ApprovedExecution{AI: bound, MaxCallDuration: 100 * time.Millisecond}, mediaSession, func(context.Context) error { return nil })
if err == nil || !strings.Contains(err.Error(), "deadline") || sdkCalls.Load() != 1 || mediaSession.sent != 1 || mediaSession.rate != 16000 {
t.Fatalf("approved opening must be played once before timed capture: requests=%d sent=%d rate=%d err=%v", sdkCalls.Load(), mediaSession.sent, mediaSession.rate, err)
}
}