289 lines
8.3 KiB
Go
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
|
|
}
|