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

294 lines
8.9 KiB
Go

package ai
import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/url"
"strings"
"sync"
"git.ipao.vip/rogee/doubao-speech-go"
"github.com/openai/openai-go/v3"
"github.com/openai/openai-go/v3/option"
)
// Call owns the keyword action for exactly one call. A failed or
// uncertain hangup is never attempted again for that call.
type Call struct {
bound Binding
keyword *KeywordHangup
mu sync.Mutex
turnMu sync.Mutex
openingRequested bool
openingReady bool
}
func NewCall(bound Binding, hangup func(context.Context) error) (*Call, error) {
if bound.Mode == "asr_only" && len(bound.HangupKeywords) != 0 {
return nil, errors.New("ASR-only call cannot play keyword closing remarks")
}
keyword, err := NewKeywordHangup(bound.HangupKeywords, hangup)
if err != nil {
return nil, err
}
return &Call{bound: bound, keyword: keyword}, nil
}
// Open synthesizes the approved opening at most once per call. An ambiguous
// synthesis result is visible to the caller and never replayed automatically.
func (c *Call) Open(ctx context.Context) ([]byte, error) {
if c == nil || c.bound.Mode != "full_ai" {
return nil, errors.New("ASR-only call has no opening TTS")
}
c.mu.Lock()
if c.openingRequested {
c.mu.Unlock()
return nil, errors.New("opening already requested; outcome unknown")
}
c.openingRequested = true
c.mu.Unlock()
if c.bound.Opening == "" {
c.mu.Lock()
c.openingReady = true
c.mu.Unlock()
return nil, nil
}
audio, err := c.bound.Synthesize(ctx, c.bound.Opening)
if err != nil {
slog.Error("approved opening synthesis failed", "stage", "opening")
return nil, err
}
c.mu.Lock()
c.openingReady = true
c.mu.Unlock()
return audio, nil
}
func (c *Call) HandleFinalASR(ctx context.Context, segment ASRSegment) (TurnResult, error) {
if c == nil {
return TurnResult{}, errors.New("AI call is unavailable")
}
group, err := c.keyword.Match(segment)
if err != nil || group == nil {
return TurnResult{}, err
}
turn := TurnResult{Transcript: segment.Text, Reply: group.ClosingRemark, EndedByKeyword: true}
slog.Info("approved keyword closing requested", "stage", "closing")
turn.AudioPCM16, err = c.bound.SynthesizeReply(ctx, group.ClosingRemark)
if err != nil {
slog.Error("approved keyword closing synthesis failed", "stage", "closing")
return turn, fmt.Errorf("keyword closing TTS failed: %w", err)
}
return turn, nil
}
func (c *Call) FinishKeyword(ctx context.Context) error {
if c == nil {
return errors.New("AI call is unavailable")
}
return c.keyword.Finish(ctx)
}
// RunTurn performs one utterance. A keyword match bypasses the LLM and
// synthesizes only the approved closing; subsequent turns are rejected.
func (c *Call) RunTurn(ctx context.Context, pcm16 []byte) (TurnResult, error) {
if c == nil {
return TurnResult{}, errors.New("AI call is unavailable")
}
c.turnMu.Lock()
defer c.turnMu.Unlock()
if c.keyword.Requested() {
return TurnResult{}, errors.New("keyword closing already requested; no further turns permitted")
}
if c.bound.Mode == "full_ai" && c.bound.Opening != "" {
c.mu.Lock()
ready := c.openingReady
c.mu.Unlock()
if !ready {
return TurnResult{}, errors.New("approved opening was not completed")
}
}
text, err := c.bound.Recognize(ctx, pcm16)
if err != nil {
return TurnResult{}, fmt.Errorf("ASR failed: %w", err)
}
closing, err := c.HandleFinalASR(ctx, ASRSegment{Source: "user", Text: text, Final: true})
if err != nil || closing.EndedByKeyword {
return closing, err
}
if c.bound.Mode == "asr_only" {
return TurnResult{Transcript: text}, nil
}
reply, err := c.bound.Complete(ctx, text)
if err != nil {
return TurnResult{}, fmt.Errorf("LLM failed: %w", err)
}
audio, err := c.bound.SynthesizeReply(ctx, reply)
if err != nil {
return TurnResult{}, fmt.Errorf("TTS failed: %w", err)
}
return TurnResult{Transcript: text, Reply: reply, AudioPCM16: audio}, nil
}
func (b Binding) Recognize(ctx context.Context, pcm16 []byte) (string, error) {
if len(pcm16) == 0 || len(pcm16)%2 != 0 {
return "", errors.New("ASR requires nonempty signed 16-bit PCM")
}
if b.ASR.Provider.Code == "ali_bailian" {
return recognizeBailian(ctx, b.ASR, pcm16)
}
u, err := url.Parse(b.ASR.Provider.Endpoint)
if err != nil || u.Host == "" || (u.Scheme != "wss" && u.Scheme != "ws") {
return "", errors.New("ASR requires a WebSocket endpoint")
}
client := doubaospeech.NewClient("",
doubaospeech.WithAPIKey(b.ASR.Provider.Credential),
doubaospeech.WithResourceID(doubaospeech.ResourceASRStreamV2),
doubaospeech.WithWebSocketURL(u.String()),
)
request := doubaospeech.ASRV2Config{
Format: doubaospeech.FormatPCMS16LE, SampleRate: 16000, Channel: 1, Bits: 16,
Request: &doubaospeech.ASRV2RequestConfig{ModelName: b.ASR.Model},
Parameters: b.ASR.Params,
}
session, err := client.ASRV2.OpenStreamSession(ctx, &request)
if err != nil {
return "", err
}
defer session.Close()
if err := session.SendAudio(ctx, pcm16, true); err != nil {
return "", err
}
var final finalASRAccumulator
for result, recvErr := range session.Recv() {
if recvErr != nil {
return "", recvErr
}
if result == nil {
continue
}
final.Add(result.Text, result.IsFinal)
if result.IsFinal {
break
}
}
return final.Result()
}
func (b Binding) Complete(ctx context.Context, finalUserText string) (string, error) {
if b.Mode != "full_ai" || b.LLM == nil {
return "", errors.New("ASR-only mode does not call LLM")
}
if strings.TrimSpace(finalUserText) == "" {
return "", errors.New("LLM requires final user text")
}
client := openai.NewClient(
option.WithAPIKey(b.LLM.Provider.Credential),
option.WithBaseURL(strings.TrimRight(b.LLM.Provider.Endpoint, "/")),
option.WithMaxRetries(0),
)
body := cloneParams(b.LLM.Params)
body["model"], _ = json.Marshal(b.LLM.Model)
body["messages"], _ = json.Marshal([]map[string]string{{"role": "system", "content": b.Prompt}, {"role": "user", "content": finalUserText}})
raw, err := json.Marshal(body)
if err != nil {
return "", errors.New("encode LLM request")
}
result, err := client.Chat.Completions.New(ctx, openai.ChatCompletionNewParams{}, option.WithRequestBody("application/json", raw))
if err != nil {
return "", err
}
if len(result.Choices) == 0 || strings.TrimSpace(result.Choices[0].Message.Content) == "" {
return "", errors.New("LLM returned no reply")
}
return result.Choices[0].Message.Content, nil
}
func (b Binding) Synthesize(ctx context.Context, text string) ([]byte, error) {
audio, _, err := b.synthesize(ctx, text, 0)
return audio, err
}
// SynthesizeReply splits a reply at the approved Unicode character limit,
// retaining every character and the same pending-audio bound across requests.
func (b Binding) SynthesizeReply(ctx context.Context, text string) ([]byte, error) {
if strings.TrimSpace(text) == "" {
return nil, errors.New("TTS requires nonempty reply")
}
maxChars := b.Conversation.SentenceMaxChars
if maxChars <= 0 {
return b.Synthesize(ctx, text)
}
runes := []rune(text)
var audio []byte
pending := 0
for start := 0; start < len(runes); start += maxChars {
end := min(start+maxChars, len(runes))
part, count, err := b.synthesize(ctx, string(runes[start:end]), pending)
if err != nil {
return nil, err
}
pending += count
audio = append(audio, part...)
}
return audio, nil
}
func (b Binding) synthesize(ctx context.Context, text string, alreadyPending int) ([]byte, int, error) {
if b.Mode != "full_ai" || b.TTS == nil {
return nil, 0, errors.New("ASR-only mode does not call TTS")
}
if strings.TrimSpace(text) == "" {
return nil, 0, errors.New("TTS requires nonempty reply")
}
limit := b.Conversation.MaxPendingAudioChunks
if limit > 0 && alreadyPending >= limit {
return nil, 0, errors.New("TTS exceeded approved pending audio chunk limit")
}
// Bailian produces one complete response for this approved sentence; a
// missing or failed response never counts as queued audio.
var audio []byte
var err error
switch b.TTS.Protocol {
case TTSProtocolDashScopeTask:
audio, err = synthesizeBailianTaskTTS(ctx, *b.TTS, text)
default:
return nil, 0, errors.New("TTS protocol is missing or unsupported")
}
if err != nil {
kind, status := bailianFailureSummary(err)
slog.Error("approved TTS synthesis failed", "category", kind, "http_status", status, "cause_type", fmt.Sprintf("%T", err))
return nil, 0, err
}
return audio, 1, nil
}
func cloneParams(params map[string]json.RawMessage) map[string]json.RawMessage {
out := make(map[string]json.RawMessage, len(params))
for key, value := range params {
out[key] = append(json.RawMessage(nil), value...)
}
return out
}
type finalASRAccumulator struct {
text string
final bool
}
func (a *finalASRAccumulator) Add(text string, final bool) {
if !a.final && final {
a.text, a.final = text, true
}
}
func (a *finalASRAccumulator) Result() (string, error) {
if !a.final || strings.TrimSpace(a.text) == "" {
return "", errors.New("ASR returned no final user transcript")
}
return a.text, nil
}