refactor(ai): retire independent snapshot and provider pipeline

This commit is contained in:
2026-09-30 14:43:18 +08:00
parent 4ef20d69c4
commit 75b415deff
18 changed files with 53 additions and 1432 deletions
-77
View File
@@ -1,77 +0,0 @@
package ai
import (
"encoding/json"
"errors"
"fmt"
"time"
"git.ipao.vip/rogee/go-sip/internal/contract"
)
type Authorization struct {
AuthorizationID string `json:"authorization_id"`
TenantID string `json:"tenant_id"`
TenantKey string `json:"tenant_key"`
AgentVersionID string `json:"agent_version_id"`
ConfigSHA256 string `json:"config_sha256"`
Mode Mode `json:"mode"`
IssuedAt string `json:"issued_at"`
ExpiresAt string `json:"expires_at"`
Source string `json:"source"`
CredentialRefs map[string]string `json:"credential_refs"`
Revoked bool `json:"revoked"`
RevocationReason string `json:"revocation_reason"`
}
// DecodeBoundAuthorization validates identity and immutable configuration binding.
// It deliberately preserves revoked/expired grants as facts, not permissions.
func DecodeBoundAuthorization(raw []byte, snapshot Snapshot, tenantID, tenantKey string) (Authorization, error) {
if err := contract.ValidateLocalAIAuthorization(raw); err != nil {
return Authorization{}, err
}
var authorization Authorization
if err := json.Unmarshal(raw, &authorization); err != nil {
return Authorization{}, err
}
if authorization.TenantID != tenantID || authorization.TenantKey != tenantKey {
return Authorization{}, errors.New("AI authorization tenant binding mismatch")
}
if authorization.AgentVersionID != snapshot.AgentVersionID || authorization.ConfigSHA256 != snapshot.Digest || authorization.Mode != snapshot.Mode {
return Authorization{}, errors.New("AI authorization does not match immutable snapshot")
}
issuedAt, err := time.Parse(time.RFC3339, authorization.IssuedAt)
if err != nil {
return Authorization{}, fmt.Errorf("parse AI authorization issued_at: %w", err)
}
expiresAt, err := time.Parse(time.RFC3339, authorization.ExpiresAt)
if err != nil {
return Authorization{}, fmt.Errorf("parse AI authorization expires_at: %w", err)
}
if !issuedAt.Before(expiresAt) {
return Authorization{}, errors.New("AI authorization has an invalid validity window")
}
return authorization, nil
}
func ValidateAuthorization(raw []byte, snapshot Snapshot, tenantID, tenantKey string, now time.Time) (Authorization, error) {
authorization, err := DecodeBoundAuthorization(raw, snapshot, tenantID, tenantKey)
if err != nil {
return Authorization{}, err
}
if authorization.Revoked {
return Authorization{}, errors.New("AI authorization is revoked")
}
issuedAt, err := time.Parse(time.RFC3339, authorization.IssuedAt)
if err != nil {
return Authorization{}, err
}
expiresAt, err := time.Parse(time.RFC3339, authorization.ExpiresAt)
if err != nil {
return Authorization{}, err
}
if now.Before(issuedAt) || !now.Before(expiresAt) {
return Authorization{}, errors.New("AI authorization is outside its validity window")
}
return authorization, nil
}
-54
View File
@@ -1,54 +0,0 @@
package ai
import (
"bytes"
"testing"
"time"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestValidateAuthorizationBindsSnapshotAndTenant(t *testing.T) {
snapshotRaw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(snapshotRaw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
authorizationRaw, err := contracts.Files.ReadFile("local/v0.3/examples/ai-authorization-v0.2.json")
if err != nil {
t.Fatal(err)
}
if _, err := ValidateAuthorization(authorizationRaw, snapshot, "tenant-1", "tenant-demo-key", time.Date(2026, 9, 18, 0, 0, 30, 0, time.UTC)); err != nil {
t.Fatal(err)
}
if _, err := ValidateAuthorization(authorizationRaw, snapshot, "other", "tenant-demo-key", time.Date(2026, 9, 18, 0, 0, 30, 0, time.UTC)); err == nil {
t.Fatal("expected tenant mismatch")
}
if _, err := ValidateAuthorization(authorizationRaw, snapshot, "tenant-1", "tenant-demo-key", time.Date(2026, 9, 18, 0, 2, 0, 0, time.UTC)); err == nil {
t.Fatal("expected expired authorization")
}
if _, err := ValidateAuthorization(bytes.Replace(authorizationRaw, []byte(`"revoked": false`), []byte(`"revoked": true`), 1), snapshot, "tenant-1", "tenant-demo-key", time.Date(2026, 9, 18, 0, 0, 30, 0, time.UTC)); err == nil {
t.Fatal("expected revocation rejection")
}
}
func TestAuthorizationRejectsRemovedEgressPoolField(t *testing.T) {
snapshotRaw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(snapshotRaw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
old, err := contracts.Files.ReadFile("upstream/v1/examples/ai-authorization.json")
if err != nil {
t.Fatal(err)
}
if _, err := DecodeBoundAuthorization(old, snapshot, "tenant-1", "tenant-demo-key"); err == nil {
t.Fatal("legacy egress-pool authorization field unexpectedly accepted")
}
}
+33
View File
@@ -91,6 +91,39 @@ func TestCurrentLLMPassesExplicitZeroAndDoesNotRetry(t *testing.T) {
}
}
func TestCurrentLLMFailureAndEmptyChoicesAreNotRetried(t *testing.T) {
for _, tc := range []struct {
name string
code int
body string
}{
{name: "provider failure", code: http.StatusInternalServerError, body: `{"error":{"message":"injected failure","type":"server_error"}}`},
{name: "empty choices", code: http.StatusOK, body: `{"id":"mock","object":"chat.completion","created":0,"model":"example-chat","choices":[]}`},
} {
t.Run(tc.name, func(t *testing.T) {
task, providers := currentFixture(t, "full_ai")
calls := 0
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
calls++
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(tc.code)
_, _ = fmt.Fprintln(w, tc.body)
}))
defer server.Close()
provider := providers["llm-example"]
provider.Endpoint = server.URL + "/v1"
providers[provider.ProviderRef] = provider
bound, err := BindCurrent(task, providers)
if err != nil {
t.Fatal(err)
}
if reply, err := bound.Complete(context.Background(), "用户说话"); err == nil || calls != 1 || reply != "" {
t.Fatalf("LLM failure must be visible without retry: reply=%q calls=%d err=%v", reply, calls, err)
}
})
}
}
func TestCurrentKeywordHangupUsesFinalUserTextOnly(t *testing.T) {
task, providers := currentFixture(t, "full_ai")
bound, err := BindCurrent(task, providers)
-77
View File
@@ -1,77 +0,0 @@
package ai
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
)
func TestLLMSDKPreservesExplicitParametersWithoutRetry(t *testing.T) {
for _, code := range []int{http.StatusOK, http.StatusInternalServerError} {
t.Run(http.StatusText(code), func(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
if r.Method != "POST" || r.URL.Path != "/v1/chat/completions" {
t.Errorf("unexpected request %s %s", r.Method, r.URL.Path)
}
var body struct {
Model string `json:"model"`
Temperature *float64 `json:"temperature"`
MaxTokens int `json:"max_tokens"`
Messages []struct{ Role, Content string } `json:"messages"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Error(err)
}
if body.Model != "approved-model" || body.Temperature == nil || *body.Temperature != 0 || body.MaxTokens != 17 || len(body.Messages) != 2 {
t.Errorf("explicit configuration not transmitted: %+v", body)
}
if len(body.Messages) == 2 && (body.Messages[0].Role != "system" || body.Messages[0].Content != "approved system" || body.Messages[1].Content != "synthetic input") {
t.Error("message content changed")
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(code)
if code == http.StatusOK {
_, _ = w.Write([]byte(`{"id":"local","object":"chat.completion","choices":[{"index":0,"message":{"role":"assistant","content":" local reply "},"finish_reason":"stop"}]}`))
} else {
_, _ = w.Write([]byte(`{"error":{"message":"injected failure","type":"server_error"}}`))
}
}))
defer server.Close()
pipeline := &ProviderPipeline{cfg: ProviderPipelineConfig{BailianAPIKey: "isolated-test-key", BailianBaseURL: server.URL + "/v1"}}
result, err := pipeline.complete(context.Background(), "approved-model", 0, 17, "approved system", "synthetic input")
if code == http.StatusOK && (err != nil || result != "local reply") {
t.Fatalf("response %q: %v", result, err)
}
if code != http.StatusOK && err == nil {
t.Fatal("provider failure hidden")
}
if calls.Load() != 1 {
t.Fatalf("automatic request retry: %d", calls.Load())
}
})
}
}
func TestLLMSDKRejectsEmptyChoices(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Error(err)
}
if len(body["messages"].([]any)) != 1 {
t.Error("invented system message")
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"id":"local","object":"chat.completion","choices":[]}`))
}))
defer server.Close()
pipeline := &ProviderPipeline{cfg: ProviderPipelineConfig{BailianAPIKey: "isolated-test-key", BailianBaseURL: server.URL + "/v1"}}
if _, err := pipeline.complete(context.Background(), "approved-model", 0, 17, "", "synthetic input"); err == nil {
t.Fatal("missing choices reported success")
}
}
-124
View File
@@ -1,124 +0,0 @@
package ai
import (
"context"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"encoding/json"
"fmt"
)
// MockTurnInput is a test-only media boundary. TranscriptHint is explicit test
// input; the mock never claims to have recognized real audio.
type MockTurnInput struct {
Audio []byte
TranscriptHint string
}
type MockTurnResult struct {
Mode Mode
Transcript string
ResponseText string
Audio []byte
Calls []string
}
type mockAIConfig struct {
Prompt struct {
Text string `json:"text"`
} `json:"prompt"`
ASR struct {
ProviderRef string `json:"provider_ref"`
} `json:"asr"`
LLM struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
} `json:"llm"`
TTS struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
} `json:"tts"`
}
// MockPipeline is a bounded provider adapter for the same callflow used by
// mixed and real modes. It never reaches a provider or accepts production
// credentials; only the adapter changes, not the business sequence.
type MockPipeline struct {
MaxAudioBytes int
}
func (p MockPipeline) Synthesize(ctx context.Context, snapshot Snapshot, text string) ([]byte, error) {
result, err := p.Run(ctx, snapshot, MockTurnInput{Audio: []byte(text), TranscriptHint: text})
if err != nil {
return nil, err
}
return result.Audio, nil
}
func (p MockPipeline) RunTurn(ctx context.Context, snapshot Snapshot, pcm16 []byte) (TurnResult, error) {
result, err := p.Run(ctx, snapshot, MockTurnInput{Audio: pcm16})
if err != nil {
return TurnResult{}, err
}
return TurnResult{Transcript: result.Transcript, Reply: result.ResponseText, AudioPCM16: result.Audio}, nil
}
func (p MockPipeline) Run(ctx context.Context, snapshot Snapshot, input MockTurnInput) (MockTurnResult, error) {
if err := ctx.Err(); err != nil {
return MockTurnResult{}, err
}
if p.MaxAudioBytes > 0 && len(input.Audio) > p.MaxAudioBytes {
return MockTurnResult{}, fmt.Errorf("mock audio exceeds limit: %d > %d", len(input.Audio), p.MaxAudioBytes)
}
var config mockAIConfig
if err := json.Unmarshal(snapshot.Raw, &config); err != nil {
return MockTurnResult{}, fmt.Errorf("decode immutable AI snapshot: %w", err)
}
transcript := input.TranscriptHint
if transcript == "" {
digest := sha256.Sum256(input.Audio)
transcript = "mock transcript " + hex.EncodeToString(digest[:4])
}
result := MockTurnResult{Mode: snapshot.Mode, Transcript: transcript, Calls: []string{"asr"}}
if snapshot.Mode == ModeASROnly {
return result, nil
}
if snapshot.Mode != ModeFullAI {
return MockTurnResult{}, fmt.Errorf("unsupported mock mode %q", snapshot.Mode)
}
if config.LLM.ProviderRef == "" || config.TTS.ProviderRef == "" || config.Prompt.Text == "" {
return MockTurnResult{}, fmt.Errorf("full-AI mock snapshot is missing provider or prompt parameters")
}
if err := ctx.Err(); err != nil {
return MockTurnResult{}, err
}
result.ResponseText = "mock response: " + transcript
result.Calls = append(result.Calls, "llm")
if err := ctx.Err(); err != nil {
return MockTurnResult{}, err
}
result.Audio = mockTonePCM16()
result.Calls = append(result.Calls, "tts")
return result, nil
}
// mockTonePCM16 is a deterministic, non-silent fixture so the isolated SIP
// callee can advance its scripted conversation after the opening prompt.
func mockTonePCM16() []byte {
const (
sampleRate = 16000
samples = sampleRate / 5
amplitude = int16(6000)
period = 36
)
pcm := make([]byte, samples*2)
for i := 0; i < samples; i++ {
value := amplitude
if (i/period)%2 == 1 {
value = -amplitude
}
binary.LittleEndian.PutUint16(pcm[i*2:i*2+2], uint16(value))
}
return pcm
}
-76
View File
@@ -1,76 +0,0 @@
package ai
import (
"context"
"testing"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestMockPipelineASROnlyStopsAfterTranscript(t *testing.T) {
raw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
result, err := (MockPipeline{MaxAudioBytes: 1024}).Run(context.Background(), snapshot, MockTurnInput{Audio: []byte{1, 2}, TranscriptHint: "客户文本"})
if err != nil {
t.Fatal(err)
}
if result.Mode != ModeASROnly || result.Transcript != "客户文本" || result.ResponseText != "" || len(result.Audio) != 0 {
t.Fatalf("unexpected ASR-only result: %+v", result)
}
if len(result.Calls) != 1 || result.Calls[0] != "asr" {
t.Fatalf("unexpected provider calls: %v", result.Calls)
}
}
func TestMockPipelineFullAIUsesAllStages(t *testing.T) {
raw, err := contracts.Read("examples/agent-version-full-explicit.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeFullAI)
if err != nil {
t.Fatal(err)
}
result, err := (MockPipeline{MaxAudioBytes: 1024}).Run(context.Background(), snapshot, MockTurnInput{TranscriptHint: "你好"})
if err != nil {
t.Fatal(err)
}
if result.ResponseText == "" || len(result.Audio) == 0 || len(result.Calls) != 3 {
t.Fatalf("unexpected full-AI result: %+v", result)
}
var nonSilent bool
for _, sample := range result.Audio {
if sample != 0 {
nonSilent = true
break
}
}
if !nonSilent {
t.Fatal("mock full-AI audio must be non-silent for isolated media tests")
}
}
func TestMockPipelineHonorsCancellationAndBound(t *testing.T) {
raw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
if _, err := (MockPipeline{MaxAudioBytes: 1024}).Run(ctx, snapshot, MockTurnInput{}); err == nil {
t.Fatal("expected cancellation")
}
if _, err := (MockPipeline{MaxAudioBytes: 1}).Run(context.Background(), snapshot, MockTurnInput{Audio: []byte{1, 2}}); err == nil {
t.Fatal("expected bounded audio rejection")
}
}
-74
View File
@@ -1,74 +0,0 @@
package ai
import (
"encoding/json"
"errors"
"fmt"
"time"
"git.ipao.vip/rogee/go-sip/internal/contract"
)
// ValidateConfigResponse binds an incoming configuration to the original,
// persisted request. It does not infer authorization from a published version.
func ValidateConfigResponse(raw []byte, request contract.ServiceMessage, now time.Time) (Snapshot, error) {
if request.MessageType != "ai.config.request" {
return Snapshot{}, errors.New("expected original AI configuration request")
}
response, err := contract.DecodeService(raw)
if err != nil {
return Snapshot{}, err
}
if response.MessageType != "ai.config.result" || response.CorrelationID != request.MessageID || response.DispatcherID != request.DispatcherID || response.TenantID != request.TenantID || response.TenantKey != request.TenantKey {
return Snapshot{}, errors.New("AI configuration response binding mismatch")
}
deadline, err := time.Parse(time.RFC3339Nano, request.NotAfter)
if err != nil {
return Snapshot{}, fmt.Errorf("AI request deadline: %w", err)
}
issued, err := time.Parse(time.RFC3339Nano, request.IssuedAt)
if err != nil {
return Snapshot{}, err
}
responseTime, err := time.Parse(time.RFC3339Nano, response.IssuedAt)
if err != nil {
return Snapshot{}, err
}
if !now.Before(deadline) || now.Before(issued) || responseTime.Before(issued) || responseTime.After(now) {
return Snapshot{}, errors.New("AI configuration response is outside the request window")
}
if response.Status != "ok" {
return Snapshot{}, fmt.Errorf("AI configuration request rejected: %s", response.ReasonCode)
}
var requested struct {
AgentVersionID string `json:"agent_version_id"`
}
if err := json.Unmarshal(request.Payload, &requested); err != nil {
return Snapshot{}, err
}
var payload struct {
Snapshot struct {
AgentVersionID string `json:"agent_version_id"`
Status string `json:"status"`
TenantID string `json:"tenant_id"`
Immutable bool `json:"immutable"`
ContentSHA256 string `json:"content_sha256"`
Config json.RawMessage `json:"config"`
} `json:"snapshot"`
}
if err := json.Unmarshal(response.Payload, &payload); err != nil {
return Snapshot{}, err
}
if (payload.Snapshot.Status != "published" && payload.Snapshot.Status != "reused") || payload.Snapshot.AgentVersionID != requested.AgentVersionID || payload.Snapshot.TenantID != request.TenantID || !payload.Snapshot.Immutable {
return Snapshot{}, errors.New("AI snapshot publication, version, tenant or immutability mismatch")
}
snapshot, err := Validate(payload.Snapshot.Config)
if err != nil {
return Snapshot{}, err
}
if snapshot.AgentVersionID != requested.AgentVersionID || snapshot.Digest != payload.Snapshot.ContentSHA256 {
return Snapshot{}, errors.New("AI configuration version or canonical digest mismatch")
}
snapshot.TenantKey = request.TenantKey
return snapshot, nil
}
-71
View File
@@ -1,71 +0,0 @@
package ai
import (
"encoding/json"
"testing"
"time"
"git.ipao.vip/rogee/go-sip/contracts"
"git.ipao.vip/rogee/go-sip/internal/contract"
)
func TestMQConfigurationMatchesOriginalRequestAndCanonicalDigest(t *testing.T) {
base := "upstream/" + contract.MQSourceCommit + "/examples/"
requestRaw, err := contracts.Files.ReadFile(base + "ai-config-request.json")
if err != nil {
t.Fatal(err)
}
request, err := contract.DecodeService(requestRaw)
if err != nil {
t.Fatal(err)
}
raw, err := contracts.Files.ReadFile(base + "ai-config-result.json")
if err != nil {
t.Fatal(err)
}
now := time.Date(2026, 9, 21, 0, 0, 1, 0, time.UTC)
snapshot, err := ValidateConfigResponse(raw, request, now)
if err != nil {
t.Fatal(err)
}
if snapshot.TenantKey != request.TenantKey || snapshot.AgentVersionID == "" {
t.Fatal("configuration lost request binding")
}
for _, field := range []string{"dispatcher_id", "tenant_id", "tenant_key", "correlation_id"} {
var changed map[string]any
if err := json.Unmarshal(raw, &changed); err != nil {
t.Fatal(err)
}
changed[field] = "incorrect"
encoded, err := json.Marshal(changed)
if err != nil {
t.Fatal(err)
}
if _, err := ValidateConfigResponse(encoded, request, now); err == nil {
t.Fatalf("mismatched %s accepted", field)
}
}
var changed map[string]any
if err := json.Unmarshal(raw, &changed); err != nil {
t.Fatal(err)
}
changed["payload"].(map[string]any)["snapshot"].(map[string]any)["status"] = "reused"
reused, err := json.Marshal(changed)
if err != nil {
t.Fatal(err)
}
if _, err := ValidateConfigResponse(reused, request, now); err != nil {
t.Fatalf("published immutable version reuse rejected: %v", err)
}
changed["payload"].(map[string]any)["snapshot"].(map[string]any)["content_sha256"] = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"
encoded, err := json.Marshal(changed)
if err != nil {
t.Fatal(err)
}
if _, err := ValidateConfigResponse(encoded, request, now); err == nil {
t.Fatal("forged digest accepted")
}
if _, err := ValidateConfigResponse(raw, request, now.Add(time.Hour)); err == nil {
t.Fatal("late response accepted")
}
}
-12
View File
@@ -1,12 +0,0 @@
package ai
import "context"
const InvalidCallMarker = "[INVALID_CALL]"
// Pipeline is the single AI turn boundary used by every call mode. The call
// flow does not know whether the implementation is mock, mixed, or real.
type Pipeline interface {
Synthesize(ctx context.Context, snapshot Snapshot, text string) ([]byte, error)
RunTurn(ctx context.Context, snapshot Snapshot, pcm16 []byte) (TurnResult, error)
}
-443
View File
@@ -1,443 +0,0 @@
package ai
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strings"
"time"
"git.ipao.vip/rogee/go-sip/internal/audio"
doubaospeech "github.com/GizClaw/doubao-speech-go"
"github.com/openai/openai-go/v3"
"github.com/openai/openai-go/v3/option"
"github.com/openai/openai-go/v3/packages/param"
"github.com/openai/openai-go/v3/shared"
)
const (
defaultBailianTTSModel = "qwen3-tts-flash"
defaultBailianTTSVoice = "Cherry"
defaultBailianLLMModel = "qwen-plus"
defaultMaxAudioBytes = 16 << 20
)
// ProviderPipelineConfig contains provider credentials and infrastructure endpoints.
// It deliberately contains no business parameters; those come from the immutable
// AI snapshot passed to RunTurn.
type ProviderPipelineConfig struct {
VolcAppID string
VolcAPIKey string
VolcWebsocketURL string
BailianAPIKey string
BailianBaseURL string
BailianTTSVoice string
HTTPClient *http.Client
MaxAudioBytes int64
}
// LoadProviderPipelineConfigFromEnv reads only credential/endpoint variables. Values
// are never returned in errors or logs.
func LoadProviderPipelineConfigFromEnv() (ProviderPipelineConfig, error) {
cfg := ProviderPipelineConfig{
VolcAppID: os.Getenv("VOLC_ASR_APP_NAME"),
VolcAPIKey: os.Getenv("VOLC_ASR_APP_KEY"),
VolcWebsocketURL: os.Getenv("VOLC_ASR_WSS_URL"),
BailianAPIKey: os.Getenv("BAILIAN_API_KEY"),
BailianBaseURL: os.Getenv("BAILIAN_BASE_URL"),
BailianTTSVoice: os.Getenv("BAILIAN_TTS_VOICE"),
HTTPClient: http.DefaultClient,
MaxAudioBytes: defaultMaxAudioBytes,
}
if cfg.BailianTTSVoice == "" {
cfg.BailianTTSVoice = defaultBailianTTSVoice
}
if cfg.VolcAppID == "" || cfg.VolcAPIKey == "" {
return ProviderPipelineConfig{}, errors.New("VOLC_ASR_APP_NAME and VOLC_ASR_APP_KEY are required")
}
return cfg, nil
}
// ProviderPipeline performs one bounded AI turn. Full-AI runs ASR -> LLM -> TTS;
// ASR-only runs ASR and returns without invoking LLM or TTS.
type ProviderPipeline struct {
cfg ProviderPipelineConfig
recognizeFn func(context.Context, string, string, []byte) (string, error)
}
func NewProviderPipeline(cfg ProviderPipelineConfig) (*ProviderPipeline, error) {
if cfg.VolcAppID == "" || cfg.VolcAPIKey == "" {
return nil, errors.New("Volcengine ASR credentials are required")
}
if cfg.BailianTTSVoice == "" {
cfg.BailianTTSVoice = defaultBailianTTSVoice
}
if cfg.HTTPClient == nil {
cfg.HTTPClient = http.DefaultClient
}
if cfg.MaxAudioBytes <= 0 {
cfg.MaxAudioBytes = defaultMaxAudioBytes
}
return &ProviderPipeline{cfg: cfg}, nil
}
type providerSnapshotConfig struct {
Mode Mode `json:"mode"`
Prompt struct {
Text string `json:"text"`
} `json:"prompt"`
ASR struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
Language string `json:"language"`
TimeoutMS int `json:"timeout_ms"`
} `json:"asr"`
LLM struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
Temperature float64 `json:"temperature"`
MaxTokens int `json:"max_tokens"`
TimeoutMS int `json:"timeout_ms"`
} `json:"llm"`
TTS struct {
ProviderRef string `json:"provider_ref"`
Model string `json:"model"`
Voice string `json:"voice"`
Speed float64 `json:"speed"`
Format json.RawMessage `json:"format"`
TimeoutMS int `json:"timeout_ms"`
} `json:"tts"`
}
// TurnResult contains only bounded facts and audio bytes needed by the caller.
// Callers must persist a hash/length, not the transcript or prompt.
type TurnResult struct {
Transcript string
Reply string
AudioPCM16 []byte
EndedByKeyword bool
InvalidCall bool
InvalidReason string
}
// Synthesize turns a bounded reply into signed 16-bit little-endian 16kHz
// mono PCM using the immutable TTS section of the snapshot.
func (p *ProviderPipeline) Synthesize(ctx context.Context, snapshot Snapshot, text string) ([]byte, error) {
if snapshot.Mode == ModeASROnly {
return nil, errors.New("ASR-only mode does not use TTS")
}
if snapshot.Mode != ModeFullAI {
return nil, fmt.Errorf("unsupported AI mode %q", snapshot.Mode)
}
var cfg providerSnapshotConfig
if err := json.Unmarshal(snapshot.Raw, &cfg); err != nil {
return nil, fmt.Errorf("decode immutable AI snapshot: %w", err)
}
if cfg.TTS.ProviderRef == "" {
return nil, errors.New("AI snapshot TTS provider ref is required")
}
if err := p.requireBailian(); err != nil {
return nil, err
}
return p.synthesize(ctx, cfg.TTS.Model, cfg.TTS.Voice, text)
}
// RunTurn executes one real AI turn from signed 16-bit little-endian 16kHz
// mono PCM. Full-AI continues through LLM/TTS; ASR-only returns after ASR. It
// is intentionally non-streaming at the provider boundary: the Asterisk media
// runtime can bound one utterance, then play returned PCM when present.
func (p *ProviderPipeline) RunTurn(ctx context.Context, snapshot Snapshot, pcm16 []byte) (TurnResult, error) {
if snapshot.Mode != ModeFullAI && snapshot.Mode != ModeASROnly {
return TurnResult{}, fmt.Errorf("unsupported AI mode %q", snapshot.Mode)
}
if len(pcm16) == 0 {
return TurnResult{}, errors.New("input PCM is empty")
}
var cfg providerSnapshotConfig
if err := json.Unmarshal(snapshot.Raw, &cfg); err != nil {
return TurnResult{}, fmt.Errorf("decode immutable AI snapshot: %w", err)
}
if cfg.ASR.ProviderRef == "" {
return TurnResult{}, errors.New("AI snapshot ASR provider ref is required")
}
if snapshot.Mode == ModeFullAI && (cfg.LLM.ProviderRef == "" || cfg.TTS.ProviderRef == "") {
return TurnResult{}, errors.New("full-AI snapshot LLM and TTS provider refs are required")
}
asrCtx := ctx
if cfg.ASR.TimeoutMS > 0 {
var cancel context.CancelFunc
asrCtx, cancel = context.WithTimeout(ctx, time.Duration(cfg.ASR.TimeoutMS)*time.Millisecond)
defer cancel()
}
recognize := p.recognize
if p.recognizeFn != nil {
recognize = p.recognizeFn
}
transcript, err := recognize(asrCtx, cfg.ASR.Model, cfg.ASR.Language, pcm16)
if err != nil {
return TurnResult{}, fmt.Errorf("ASR failed: %w", err)
}
if strings.TrimSpace(transcript) == "" {
return TurnResult{}, errors.New("ASR returned empty transcript")
}
if snapshot.Mode == ModeASROnly {
return TurnResult{Transcript: transcript}, nil
}
if err := p.requireBailian(); err != nil {
return TurnResult{}, err
}
llmCtx := ctx
if cfg.LLM.TimeoutMS > 0 {
var cancel context.CancelFunc
llmCtx, cancel = context.WithTimeout(ctx, time.Duration(cfg.LLM.TimeoutMS)*time.Millisecond)
defer cancel()
}
reply, err := p.complete(llmCtx, cfg.LLM.Model, cfg.LLM.Temperature, cfg.LLM.MaxTokens, cfg.Prompt.Text, transcript)
if err != nil {
return TurnResult{}, fmt.Errorf("LLM failed: %w", err)
}
if strings.TrimSpace(reply) == "" {
return TurnResult{}, errors.New("LLM returned empty reply")
}
ttsCtx := ctx
if cfg.TTS.TimeoutMS > 0 {
var cancel context.CancelFunc
ttsCtx, cancel = context.WithTimeout(ctx, time.Duration(cfg.TTS.TimeoutMS)*time.Millisecond)
defer cancel()
}
audio, err := p.synthesize(ttsCtx, cfg.TTS.Model, cfg.TTS.Voice, reply)
if err != nil {
return TurnResult{}, fmt.Errorf("TTS failed: %w", err)
}
return TurnResult{Transcript: transcript, Reply: reply, AudioPCM16: audio}, nil
}
func (p *ProviderPipeline) requireBailian() error {
if p.cfg.BailianAPIKey == "" || p.cfg.BailianBaseURL == "" {
return errors.New("Bailian credentials and base URL are required for full-AI mode")
}
return nil
}
func (p *ProviderPipeline) recognize(ctx context.Context, model, language string, pcm16 []byte) (string, error) {
client := doubaospeech.NewClient(p.cfg.VolcAppID,
doubaospeech.WithAPIKey(p.cfg.VolcAPIKey),
doubaospeech.WithWebSocketURL(p.cfg.VolcWebsocketURL),
)
if p.cfg.VolcWebsocketURL == "" {
client = doubaospeech.NewClient(p.cfg.VolcAppID, doubaospeech.WithAPIKey(p.cfg.VolcAPIKey))
}
lang := doubaospeech.LanguageZhCN
if language != "" {
lang = doubaospeech.Language(language)
}
request := &doubaospeech.ASRV2RequestConfig{
ModelName: model,
ResultType: "single",
EnableITN: boolPtr(true),
EnablePunc: boolPtr(true),
EnableNonstream: boolPtr(true),
}
session, err := client.ASRV2.OpenStreamSession(ctx, &doubaospeech.ASRV2Config{
Format: doubaospeech.FormatPCM,
SampleRate: doubaospeech.SampleRate16000,
Channel: 1,
Bits: 16,
Language: lang,
Request: request,
ResultType: "single",
})
if err != nil {
return "", err
}
defer session.Close()
if err := session.SendAudio(ctx, pcm16, true); err != nil {
return "", err
}
var transcript string
for result, recvErr := range session.Recv() {
if recvErr != nil {
return "", recvErr
}
if result != nil && result.Text != "" {
transcript = result.Text
}
if result != nil && result.IsFinal {
break
}
}
return transcript, nil
}
func (p *ProviderPipeline) complete(ctx context.Context, model string, temperature float64, maxTokens int, systemPrompt, transcript string) (string, error) {
if model == "" {
model = defaultBailianLLMModel
}
client := openai.NewClient(
option.WithAPIKey(p.cfg.BailianAPIKey),
option.WithBaseURL(strings.TrimRight(p.cfg.BailianBaseURL, "/")),
option.WithMaxRetries(0),
)
messages := make([]openai.ChatCompletionMessageParamUnion, 0, 2)
if strings.TrimSpace(systemPrompt) != "" {
messages = append(messages, openai.SystemMessage(systemPrompt))
}
messages = append(messages, openai.UserMessage(transcript))
params := openai.ChatCompletionNewParams{
Model: shared.ChatModel(model),
Messages: messages,
}
if maxTokens > 0 {
params.MaxTokens = param.NewOpt(int64(maxTokens))
}
if temperature >= 0 {
params.Temperature = param.NewOpt(temperature)
}
result, err := client.Chat.Completions.New(ctx, params)
if err != nil {
return "", err
}
if len(result.Choices) == 0 {
return "", errors.New("LLM response has no choices")
}
return strings.TrimSpace(result.Choices[0].Message.Content), nil
}
func (p *ProviderPipeline) synthesize(ctx context.Context, model, voice, text string) ([]byte, error) {
if model == "" {
model = defaultBailianTTSModel
}
if strings.TrimSpace(voice) == "" || strings.HasPrefix(voice, "env:") {
voice = p.cfg.BailianTTSVoice
}
if voice == "" {
voice = defaultBailianTTSVoice
}
base := p.cfg.BailianBaseURL
u, err := url.Parse(base)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, errors.New("invalid BAILIAN_BASE_URL")
}
generationURL := (&url.URL{Scheme: u.Scheme, Host: u.Host, Path: "/api/v1/services/aigc/multimodal-generation/generation"}).String()
body, err := json.Marshal(map[string]any{
"model": model,
"input": map[string]any{
"text": text,
"voice": voice,
"language_type": "Chinese",
},
})
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, generationURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+p.cfg.BailianAPIKey)
req.Header.Set("Content-Type", "application/json")
resp, err := p.cfg.HTTPClient.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1024))
return nil, fmt.Errorf("TTS generation HTTP %d", resp.StatusCode)
}
var envelope struct {
Output struct {
Audio struct {
URL string `json:"url"`
} `json:"audio"`
} `json:"output"`
}
if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&envelope); err != nil {
return nil, err
}
if envelope.Output.Audio.URL == "" {
return nil, errors.New("TTS response has no audio URL")
}
audioReq, err := http.NewRequestWithContext(ctx, http.MethodGet, envelope.Output.Audio.URL, nil)
if err != nil {
return nil, err
}
audioResp, err := p.cfg.HTTPClient.Do(audioReq)
if err != nil {
return nil, err
}
defer audioResp.Body.Close()
if audioResp.StatusCode/100 != 2 {
return nil, fmt.Errorf("TTS audio URL HTTP %d", audioResp.StatusCode)
}
wav, err := io.ReadAll(io.LimitReader(audioResp.Body, p.cfg.MaxAudioBytes+1))
if err != nil {
return nil, err
}
if int64(len(wav)) > p.cfg.MaxAudioBytes {
return nil, errors.New("TTS audio exceeds configured limit")
}
return decodeWAVToPCM16(wav)
}
func boolPtr(value bool) *bool { return &value }
func decodeWAVToPCM16(data []byte) ([]byte, error) {
if len(data) < 12 || string(data[:4]) != "RIFF" || string(data[8:12]) != "WAVE" {
return nil, errors.New("TTS audio is not RIFF/WAVE")
}
var format, channels, sampleRate, bits int
var pcm []byte
for pos := 12; pos+8 <= len(data); {
id := string(data[pos : pos+4])
size := int(binary.LittleEndian.Uint32(data[pos+4 : pos+8]))
pos += 8
if size < 0 || pos+size > len(data) {
// DashScope's streaming WAV uses a 0x7fffffff placeholder for
// RIFF/data sizes. The response body is authoritative and bounded
// by MaxAudioBytes, so only the data chunk may consume the remainder.
if id != "data" {
return nil, errors.New("invalid WAV chunk size")
}
size = len(data) - pos
}
switch id {
case "fmt ":
if size < 16 {
return nil, errors.New("invalid WAV fmt chunk")
}
format = int(binary.LittleEndian.Uint16(data[pos : pos+2]))
channels = int(binary.LittleEndian.Uint16(data[pos+2 : pos+4]))
sampleRate = int(binary.LittleEndian.Uint32(data[pos+4 : pos+8]))
bits = int(binary.LittleEndian.Uint16(data[pos+14 : pos+16]))
case "data":
pcm = append([]byte(nil), data[pos:pos+size]...)
}
pos += size
if size%2 == 1 {
pos++
}
}
if format != 1 || channels != 1 || bits != 16 || sampleRate <= 0 || len(pcm) == 0 {
return nil, fmt.Errorf("unsupported WAV format=%d channels=%d rate=%d bits=%d", format, channels, sampleRate, bits)
}
if sampleRate == 16000 {
return pcm, nil
}
return resamplePCM16(pcm, sampleRate, 16000), nil
}
func resamplePCM16(src []byte, sourceRate, targetRate int) []byte {
return audio.ResamplePCM16(src, sourceRate, targetRate)
}
-168
View File
@@ -1,168 +0,0 @@
package ai
import (
"context"
"encoding/binary"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func providerSnapshotRaw(mode string, asr, llm, tts bool) []byte {
value := map[string]any{
"mode": mode,
"prompt": map[string]any{"text": "approved system"},
"asr": map[string]any{"provider_ref": "asr-ref", "model": "asr-model", "language": "zh-CN", "timeout_ms": 50},
"llm": map[string]any{"provider_ref": "llm-ref", "model": "llm-model", "temperature": 0.0, "max_tokens": 17, "timeout_ms": 50},
"tts": map[string]any{"provider_ref": "tts-ref", "model": "tts-model", "voice": "voice-a", "speed": 1.0, "timeout_ms": 50},
}
if !asr {
delete(value["asr"].(map[string]any), "provider_ref")
}
if !llm {
delete(value["llm"].(map[string]any), "provider_ref")
}
if !tts {
delete(value["tts"].(map[string]any), "provider_ref")
}
raw, _ := json.Marshal(value)
return raw
}
func wavPCM16(sampleRate int, pcm []byte) []byte {
data := make([]byte, 44+len(pcm))
copy(data[:4], "RIFF")
binary.LittleEndian.PutUint32(data[4:8], uint32(len(data)-8))
copy(data[8:12], "WAVE")
copy(data[12:16], "fmt ")
binary.LittleEndian.PutUint32(data[16:20], 16)
binary.LittleEndian.PutUint16(data[20:22], 1)
binary.LittleEndian.PutUint16(data[22:24], 1)
binary.LittleEndian.PutUint32(data[24:28], uint32(sampleRate))
binary.LittleEndian.PutUint32(data[28:32], uint32(sampleRate*2))
binary.LittleEndian.PutUint16(data[32:34], 2)
binary.LittleEndian.PutUint16(data[34:36], 16)
copy(data[36:40], "data")
binary.LittleEndian.PutUint32(data[40:44], uint32(len(pcm)))
copy(data[44:], pcm)
return data
}
func TestProviderPipelineFullTurnUsesConfiguredLLMAndTTS(t *testing.T) {
pcm := []byte{1, 0, 2, 0, 3, 0, 4, 0}
wav := wavPCM16(16000, pcm)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch {
case strings.HasSuffix(r.URL.Path, "/chat/completions"):
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"approved reply"}}]}`))
case strings.HasSuffix(r.URL.Path, "/generation"):
_, _ = w.Write([]byte(`{"output":{"audio":{"url":"` + "PLACEHOLDER" + `"}}}`))
case r.URL.Path == "/audio.wav":
w.Header().Set("Content-Type", "audio/wav")
_, _ = w.Write(wav)
default:
http.NotFound(w, r)
}
}))
defer server.Close()
// Replace the placeholder without putting a second server or a public URL in
// the fixture; the provider still uses the same bounded test HTTP client.
server.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch {
case strings.HasSuffix(r.URL.Path, "/chat/completions"):
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"approved reply"}}]}`))
case strings.HasSuffix(r.URL.Path, "/generation"):
_, _ = w.Write([]byte(`{"output":{"audio":{"url":"` + server.URL + `/audio.wav"}}}`))
case r.URL.Path == "/audio.wav":
w.Header().Set("Content-Type", "audio/wav")
_, _ = w.Write(wav)
default:
http.NotFound(w, r)
}
})
pipeline := &ProviderPipeline{cfg: ProviderPipelineConfig{
VolcAppID: "asr-app", VolcAPIKey: "asr-key", BailianAPIKey: "llm-key",
BailianBaseURL: server.URL + "/v1", HTTPClient: server.Client(), MaxAudioBytes: 1024,
}}
pipeline.recognizeFn = func(context.Context, string, string, []byte) (string, error) { return "approved transcript", nil }
result, err := pipeline.RunTurn(context.Background(), Snapshot{Mode: ModeFullAI, Raw: providerSnapshotRaw("full_ai", true, true, true)}, []byte{1, 2})
if err != nil {
t.Fatal(err)
}
if result.Transcript != "approved transcript" || result.Reply != "approved reply" || string(result.AudioPCM16) != string(pcm) {
t.Fatalf("unexpected full turn result: %+v", result)
}
}
func TestProviderPipelineRejectsInvalidTurnInputs(t *testing.T) {
base := &ProviderPipeline{cfg: ProviderPipelineConfig{VolcAppID: "asr", VolcAPIKey: "key"}}
base.recognizeFn = func(context.Context, string, string, []byte) (string, error) { return "", nil }
cases := []struct {
name string
pipeline *ProviderPipeline
snapshot Snapshot
pcm []byte
want string
}{
{"empty pcm", base, Snapshot{Mode: ModeASROnly, Raw: providerSnapshotRaw("asr_only", true, false, false)}, nil, "input PCM is empty"},
{"invalid mode", base, Snapshot{Mode: Mode("invalid"), Raw: providerSnapshotRaw("invalid", true, false, false)}, []byte{1}, "unsupported AI mode"},
{"bad snapshot", base, Snapshot{Mode: ModeASROnly, Raw: []byte("{")}, []byte{1}, "decode immutable AI snapshot"},
{"missing asr", base, Snapshot{Mode: ModeASROnly, Raw: providerSnapshotRaw("asr_only", false, false, false)}, []byte{1}, "ASR provider ref is required"},
{"empty transcript", base, Snapshot{Mode: ModeASROnly, Raw: providerSnapshotRaw("asr_only", true, false, false)}, []byte{1}, "ASR returned empty transcript"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if _, err := tc.pipeline.RunTurn(context.Background(), tc.snapshot, tc.pcm); err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("error=%v, want %q", err, tc.want)
}
})
}
failingASR := &ProviderPipeline{cfg: base.cfg, recognizeFn: func(context.Context, string, string, []byte) (string, error) { return "", context.DeadlineExceeded }}
if _, err := failingASR.RunTurn(context.Background(), Snapshot{Mode: ModeASROnly, Raw: providerSnapshotRaw("asr_only", true, false, false)}, []byte{1}); err == nil || !strings.Contains(err.Error(), "ASR failed") {
t.Fatalf("ASR error was hidden: %v", err)
}
}
func TestProviderPipelineWAVValidationAndResampling(t *testing.T) {
pcm := []byte{1, 0, 2, 0, 3, 0, 4, 0}
for _, tc := range []struct {
name string
data []byte
want string
}{
{"short", []byte("RIFF"), "not RIFF/WAVE"},
{"wrong container", []byte("RIFFxxxxNOPE"), "not RIFF/WAVE"},
{"bad fmt", append([]byte("RIFFxxxxWAVEfmt "), 2, 0, 0, 0, 0, 0), "invalid WAV fmt"},
{"bad format", wavPCM16(16000, pcm)[:44], "unsupported WAV format"},
} {
t.Run(tc.name, func(t *testing.T) {
if _, err := decodeWAVToPCM16(tc.data); err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("error=%v, want %q", err, tc.want)
}
})
}
resampled, err := decodeWAVToPCM16(wavPCM16(8000, pcm))
if err != nil || len(resampled) == 0 {
t.Fatalf("valid non-16k WAV was not resampled: len=%d err=%v", len(resampled), err)
}
if got := resamplePCM16(nil, 8000, 16000); len(got) != 0 {
t.Fatalf("empty resample returned %d bytes", len(got))
}
}
func TestProviderPipelineSynthesizeFailureModes(t *testing.T) {
p := &ProviderPipeline{cfg: ProviderPipelineConfig{VolcAppID: "asr", VolcAPIKey: "key", BailianAPIKey: "key", BailianBaseURL: "://bad"}}
if _, err := p.Synthesize(context.Background(), Snapshot{Mode: Mode("other")}, "text"); err == nil || !strings.Contains(err.Error(), "unsupported AI mode") {
t.Fatal("unsupported mode accepted")
}
if _, err := p.Synthesize(context.Background(), Snapshot{Mode: ModeFullAI, Raw: []byte(`{"mode":"full_ai"}`)}, "text"); err == nil || !strings.Contains(err.Error(), "TTS provider ref") {
t.Fatal("missing TTS provider accepted")
}
if _, err := p.Synthesize(context.Background(), Snapshot{Mode: ModeFullAI, Raw: providerSnapshotRaw("full_ai", true, true, true)}, "text"); err == nil || !strings.Contains(err.Error(), "invalid BAILIAN_BASE_URL") {
t.Fatal("invalid TTS endpoint accepted")
}
}
@@ -1,46 +0,0 @@
package ai
import (
"context"
"os"
"testing"
"time"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestProviderFullAIChain(t *testing.T) {
if os.Getenv("AGENT_CALL_PROVIDER_SMOKE") != "1" {
t.Skip("set AGENT_CALL_PROVIDER_SMOKE=1 to authorize the provider smoke")
}
cfg, err := LoadProviderPipelineConfigFromEnv()
if err != nil {
t.Fatal(err)
}
pipeline, err := NewProviderPipeline(cfg)
if err != nil {
t.Fatal(err)
}
raw, err := contracts.Read("examples/agent-version-full-production-v1.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeFullAI)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
defer cancel()
input, err := pipeline.Synthesize(ctx, snapshot, "你好,我想了解单节点服务。")
if err != nil {
t.Fatal(err)
}
result, err := pipeline.RunTurn(ctx, snapshot, input)
if err != nil {
t.Fatal(err)
}
if result.Transcript == "" || result.Reply == "" || len(result.AudioPCM16) == 0 {
t.Fatalf("incomplete provider chain transcript=%d reply=%d audio=%d", len([]rune(result.Transcript)), len([]rune(result.Reply)), len(result.AudioPCM16))
}
t.Logf("provider full-AI chain passed transcript_chars=%d reply_chars=%d audio_bytes=%d", len([]rune(result.Transcript)), len([]rune(result.Reply)), len(result.AudioPCM16))
}
-54
View File
@@ -1,54 +0,0 @@
package ai
import (
"context"
"strings"
"testing"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestLoadProviderPipelineConfigAllowsASROnlyCredentials(t *testing.T) {
t.Setenv("VOLC_ASR_APP_NAME", "asr-app")
t.Setenv("VOLC_ASR_APP_KEY", "asr-key")
t.Setenv("VOLC_ASR_WSS_URL", "")
t.Setenv("BAILIAN_API_KEY", "")
t.Setenv("BAILIAN_BASE_URL", "")
t.Setenv("BAILIAN_TTS_VOICE", "")
if _, err := LoadProviderPipelineConfigFromEnv(); err != nil {
t.Fatalf("ASR-only provider config should not require Bailian credentials: %v", err)
}
}
func TestProviderPipelineRunsASROnlyWithoutBailian(t *testing.T) {
raw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
pipeline, err := NewProviderPipeline(ProviderPipelineConfig{VolcAppID: "asr-app", VolcAPIKey: "asr-key"})
if err != nil {
t.Fatal(err)
}
pipeline.recognizeFn = func(context.Context, string, string, []byte) (string, error) {
return "recognized without Bailian", nil
}
result, err := pipeline.RunTurn(context.Background(), snapshot, []byte{1, 2})
if err != nil {
t.Fatal(err)
}
if result.Transcript != "recognized without Bailian" || result.Reply != "" || len(result.AudioPCM16) != 0 {
t.Fatalf("unexpected ASR-only result: %+v", result)
}
}
func TestProviderPipelineASROnlyDoesNotUseTTS(t *testing.T) {
pipeline := &ProviderPipeline{}
_, err := pipeline.Synthesize(context.Background(), Snapshot{Mode: ModeASROnly}, "opening")
if err == nil || !strings.Contains(err.Error(), "does not use TTS") {
t.Fatalf("Synthesize(ASR-only) error=%v", err)
}
}
-80
View File
@@ -1,80 +0,0 @@
// Package ai handles immutable Agent-version snapshots. The JSON Schema is
// loaded from the pinned upstream contract bundle; this package does not
// duplicate or loosen that schema.
package ai
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"unicode/utf8"
"github.com/cyberphone/json-canonicalization/go/src/webpki.org/jsoncanonicalizer"
"git.ipao.vip/rogee/go-sip/contracts"
"git.ipao.vip/rogee/go-sip/internal/contract"
)
type Mode string
const (
ModeFullAI Mode = "full_ai"
ModeASROnly Mode = "asr_only"
)
type Snapshot struct {
TenantKey string
AgentVersionID string
Digest string
Raw []byte
Mode Mode
}
func Validate(raw []byte) (Snapshot, error) {
if !utf8.Valid(raw) {
return Snapshot{}, errors.New("AI configuration must be valid UTF-8")
}
if err := contract.ValidateSourceSchema("ai-config.schema.json", raw); err != nil {
return Snapshot{}, err
}
var value struct {
AgentVersionID string `json:"agent_version_id"`
Mode Mode `json:"mode"`
}
if err := json.Unmarshal(raw, &value); err != nil {
return Snapshot{}, err
}
mode := value.Mode
if mode == "" {
// Legacy immutable snapshots predate the explicit mode field and are
// full-AI snapshots because llm/prompt/tts are required by that branch.
mode = ModeFullAI
}
if mode != ModeFullAI && mode != ModeASROnly {
return Snapshot{}, fmt.Errorf("unsupported AI mode %q", mode)
}
canonical, err := jsoncanonicalizer.Transform(raw)
if err != nil {
return Snapshot{}, fmt.Errorf("canonicalize AI configuration: %w", err)
}
digest := sha256.Sum256(canonical)
return Snapshot{AgentVersionID: value.AgentVersionID, Digest: hex.EncodeToString(digest[:]), Raw: append([]byte(nil), raw...), Mode: mode}, nil
}
func ValidateForMode(raw []byte, mode Mode) (Snapshot, error) {
if mode != ModeFullAI && mode != ModeASROnly {
return Snapshot{}, fmt.Errorf("unsupported AI mode %q", mode)
}
snapshot, err := Validate(raw)
if err != nil {
return Snapshot{}, err
}
if snapshot.Mode != mode {
return Snapshot{}, fmt.Errorf("AI config mode %q does not match requested mode %q", snapshot.Mode, mode)
}
return snapshot, nil
}
func ContractSource() string { return contracts.SourceCommit }
-31
View File
@@ -1,31 +0,0 @@
package ai
import (
"bytes"
"testing"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestSnapshotDigestUsesJCSNotJSONPresentation(t *testing.T) {
raw, err := contracts.Read("examples/agent-version.json")
if err != nil {
t.Fatal(err)
}
first, err := Validate(raw)
if err != nil {
t.Fatal(err)
}
variant := bytes.ReplaceAll(raw, []byte(`"speed": 1.0`), []byte(`"speed": 1`))
variant = append([]byte("\n\t"), variant...)
second, err := Validate(variant)
if err != nil {
t.Fatal(err)
}
if first.Digest != second.Digest {
t.Fatal("JSON presentation changed immutable AI digest")
}
if !bytes.Equal(first.Raw, raw) {
t.Fatal("canonicalization changed supplied snapshot bytes")
}
}
-45
View File
@@ -1,45 +0,0 @@
package ai
import (
"testing"
"git.ipao.vip/rogee/go-sip/contracts"
)
func TestValidateUsesPinnedAIContract(t *testing.T) {
raw, err := contracts.Read("examples/agent-version.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeFullAI)
if err != nil {
t.Fatal(err)
}
if snapshot.AgentVersionID == "" || snapshot.Digest == "" || snapshot.Mode != ModeFullAI {
t.Fatalf("snapshot = %+v", snapshot)
}
}
func TestValidateASROnlySnapshot(t *testing.T) {
raw, err := contracts.Read("examples/agent-version-asr-only.json")
if err != nil {
t.Fatal(err)
}
snapshot, err := ValidateForMode(raw, ModeASROnly)
if err != nil {
t.Fatal(err)
}
if snapshot.AgentVersionID == "" || snapshot.Digest == "" || snapshot.Mode != ModeASROnly {
t.Fatalf("snapshot = %+v", snapshot)
}
}
func TestASROnlyRejectsFullAIFields(t *testing.T) {
raw, err := contracts.Read("examples/invalid-asr-only-with-llm.json")
if err != nil {
t.Fatal(err)
}
if _, err := ValidateForMode(raw, ModeASROnly); err == nil {
t.Fatal("expected ASR-only schema rejection")
}
}
+19
View File
@@ -0,0 +1,19 @@
package ai
// Mode is the approved AI execution mode bound to one call.
type Mode string
const (
ModeFullAI Mode = "full_ai"
ModeASROnly Mode = "asr_only"
)
// TurnResult contains the observed speech and any approved response or stop fact.
type TurnResult struct {
Transcript string
Reply string
AudioPCM16 []byte
EndedByKeyword bool
InvalidCall bool
InvalidReason string
}