443 lines
14 KiB
Go
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)
|
|
}
|