refactor(ai): retire independent snapshot and provider pipeline
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user