Files
go-sip/internal/ai/provider_pipeline.go
T

443 lines
14 KiB
Go

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
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)
}