294 lines
8.9 KiB
Go
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
|
|
}
|