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

289 lines
8.3 KiB
Go

package ai
import (
"context"
"errors"
"fmt"
"net/url"
"strings"
"sync"
"time"
"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"
)
// CurrentCall owns the keyword action for exactly one call. A failed or
// uncertain hangup is never attempted again for that call.
type CurrentCall struct {
bound CurrentBound
keyword *KeywordHangup
mu sync.Mutex
openingRequested bool
openingReady bool
}
func NewCurrentCall(bound CurrentBound, hangup func(context.Context) error) (*CurrentCall, error) {
keyword, err := NewKeywordHangup(bound.HangupKeywords, hangup)
if err != nil {
return nil, err
}
return &CurrentCall{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 *CurrentCall) 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 {
return nil, err
}
c.mu.Lock()
c.openingReady = true
c.mu.Unlock()
return audio, nil
}
func (c *CurrentCall) HandleFinalASR(ctx context.Context, segment ASRSegment) (bool, error) {
if c == nil {
return false, errors.New("AI call is unavailable")
}
return c.keyword.Handle(ctx, segment)
}
// RunTurn performs one utterance. No LLM or TTS request is sent after a
// keyword hangup, and only a final user recognition may enter the LLM.
func (c *CurrentCall) RunTurn(ctx context.Context, pcm16 []byte) (TurnResult, error) {
if c == nil {
return TurnResult{}, errors.New("AI call is unavailable")
}
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)
}
stopped, err := c.HandleFinalASR(ctx, ASRSegment{Source: "user", Text: text, Final: true})
if err != nil {
return TurnResult{}, err
}
if stopped {
return TurnResult{Transcript: text, EndedByKeyword: true}, nil
}
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 CurrentBound) 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")
}
ctx, cancel := currentDeadline(ctx, b.ASR.Timeout)
defer cancel()
u, err := url.Parse(b.ASR.Provider.Endpoint)
if err != nil {
return "", errors.New("ASR endpoint is invalid")
}
if u.Scheme == "https" {
u.Scheme = "wss"
} else if u.Scheme == "http" {
u.Scheme = "ws"
}
client := doubaospeech.NewClient("",
doubaospeech.WithAPIKey(b.ASR.Provider.Credential),
doubaospeech.WithWebSocketURL(u.String()),
)
request := b.ASR.Request
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 CurrentBound) 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")
}
ctx, cancel := currentDeadline(ctx, b.LLM.Timeout)
defer cancel()
client := openai.NewClient(
option.WithAPIKey(b.LLM.Provider.Credential),
option.WithBaseURL(strings.TrimRight(b.LLM.Provider.Endpoint, "/")),
option.WithMaxRetries(0),
)
messages := []openai.ChatCompletionMessageParamUnion{openai.SystemMessage(b.Prompt), openai.UserMessage(finalUserText)}
params := openai.ChatCompletionNewParams{Model: shared.ChatModel(b.LLM.Model), Messages: messages}
if b.LLM.Temperature != nil {
params.Temperature = param.NewOpt(*b.LLM.Temperature)
}
if b.LLM.MaxTokens != nil {
params.MaxTokens = param.NewOpt(*b.LLM.MaxTokens)
}
result, err := client.Chat.Completions.New(ctx, params)
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 CurrentBound) 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 CurrentBound) 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 CurrentBound) 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")
}
ctx, cancel := currentDeadline(ctx, b.TTS.Timeout)
defer cancel()
request := b.TTS.Request // per-call copy: concurrent calls never share mutable SDK parameters
request.Text = text
client := doubaospeech.NewClient("",
doubaospeech.WithAPIKey(b.TTS.Provider.Credential),
doubaospeech.WithBaseURL(b.TTS.Provider.Endpoint),
)
var audio []byte
chunks := 0
completed := false
for chunk, err := range client.TTSV2.Stream(ctx, &request) {
if err != nil {
return nil, chunks, err
}
if chunk == nil {
return nil, chunks, errors.New("TTS returned an empty stream chunk")
}
if len(chunk.Audio) > 0 {
chunks++
if limit > 0 && alreadyPending+chunks > limit {
return nil, chunks, errors.New("TTS exceeded approved pending audio chunk limit")
}
audio = append(audio, chunk.Audio...)
}
if chunk.IsLast {
completed = true
break
}
}
if !completed || len(audio) == 0 || len(audio)%2 != 0 {
return nil, chunks, errors.New("TTS stream ended without complete PCM16 audio")
}
return audio, chunks, nil
}
func currentDeadline(ctx context.Context, timeout time.Duration) (context.Context, context.CancelFunc) {
if timeout > 0 {
return context.WithTimeout(ctx, timeout)
}
return context.WithCancel(ctx)
}
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
}