112 lines
2.6 KiB
Go
112 lines
2.6 KiB
Go
package pipeline
|
|
|
|
import (
|
|
"context"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"ai-operator/internal/ai/llm"
|
|
"ai-operator/internal/ai/stt"
|
|
"ai-operator/internal/ai/tts"
|
|
"ai-operator/internal/media"
|
|
)
|
|
|
|
type FakeSTT struct {
|
|
events chan stt.Event
|
|
audio int
|
|
}
|
|
|
|
func NewFakeSTT() *FakeSTT { return &FakeSTT{events: make(chan stt.Event, 16)} }
|
|
|
|
func (f *FakeSTT) Start(ctx context.Context, req stt.StreamRequest) error { return nil }
|
|
|
|
func (f *FakeSTT) SendAudio(ctx context.Context, chunk media.AudioChunk) error {
|
|
f.audio += len(chunk.Data)
|
|
f.events <- stt.Event{Type: stt.EventPartialTranscript, Text: "Сколько стоит", At: time.Now().UTC()}
|
|
f.events <- stt.Event{Type: stt.EventCommittedTranscript, Text: "Сколько стоит первичное подключение газа?", Language: "ru", At: time.Now().UTC()}
|
|
return nil
|
|
}
|
|
|
|
func (f *FakeSTT) Events() <-chan stt.Event { return f.events }
|
|
|
|
func (f *FakeSTT) Close(ctx context.Context) error {
|
|
close(f.events)
|
|
return nil
|
|
}
|
|
|
|
type FakeLLM struct {
|
|
Text string
|
|
}
|
|
|
|
func (f FakeLLM) StreamGenerate(ctx context.Context, req llm.GenerateRequest) (<-chan llm.Event, error) {
|
|
out := make(chan llm.Event, 8)
|
|
text := f.Text
|
|
if text == "" {
|
|
text = "Первичное подключение газа к газовому оборудованию осуществляется бесплатно."
|
|
}
|
|
go func() {
|
|
defer close(out)
|
|
parts := strings.SplitAfter(text, " ")
|
|
for _, p := range parts {
|
|
if p != "" {
|
|
out <- llm.Event{Type: llm.EventTextDelta, Text: p}
|
|
}
|
|
}
|
|
out <- llm.Event{Type: llm.EventDone}
|
|
}()
|
|
return out, nil
|
|
}
|
|
|
|
type FakeTTSFactory struct {
|
|
mu sync.Mutex
|
|
Texts []string
|
|
Starts int
|
|
Cancels int
|
|
}
|
|
|
|
func (f *FakeTTSFactory) New(language string) tts.StreamingTTS {
|
|
return &FakeTTS{factory: f, audio: make(chan tts.AudioChunk, 16)}
|
|
}
|
|
|
|
type FakeTTS struct {
|
|
factory *FakeTTSFactory
|
|
audio chan tts.AudioChunk
|
|
closed bool
|
|
}
|
|
|
|
func (f *FakeTTS) Start(ctx context.Context, req tts.TTSStreamRequest) error {
|
|
f.factory.mu.Lock()
|
|
f.factory.Starts++
|
|
f.factory.mu.Unlock()
|
|
return nil
|
|
}
|
|
|
|
func (f *FakeTTS) SendText(ctx context.Context, text string, flush bool) error {
|
|
f.factory.mu.Lock()
|
|
if text != "" {
|
|
f.factory.Texts = append(f.factory.Texts, text)
|
|
}
|
|
f.factory.mu.Unlock()
|
|
if text != "" {
|
|
select {
|
|
case f.audio <- tts.AudioChunk{Data: []byte{0, 1, 0, 1}, Timestamp: time.Now().UTC()}:
|
|
default:
|
|
}
|
|
}
|
|
if flush && text == "" {
|
|
f.Close(ctx)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (f *FakeTTS) Audio() <-chan tts.AudioChunk { return f.audio }
|
|
|
|
func (f *FakeTTS) Close(ctx context.Context) error {
|
|
if !f.closed {
|
|
close(f.audio)
|
|
f.closed = true
|
|
}
|
|
return nil
|
|
}
|