662 lines
19 KiB
Go
662 lines
19 KiB
Go
package orchestrator
|
||
|
||
import (
|
||
"context"
|
||
"encoding/base64"
|
||
"errors"
|
||
"testing"
|
||
"time"
|
||
|
||
"github.com/stretchr/testify/assert"
|
||
"github.com/stretchr/testify/mock"
|
||
|
||
"github.com/hhs/camtalk/internal/ai/llm"
|
||
"github.com/hhs/camtalk/internal/ai/stt"
|
||
"github.com/hhs/camtalk/internal/ai/tts"
|
||
"github.com/hhs/camtalk/internal/logger"
|
||
"github.com/hhs/camtalk/internal/models"
|
||
)
|
||
|
||
func init() {
|
||
logger.Init("debug", "console")
|
||
}
|
||
|
||
// MockSTTService mock STT 服务
|
||
type MockSTTService struct {
|
||
mock.Mock
|
||
}
|
||
|
||
func (m *MockSTTService) Recognize(ctx context.Context, audio []byte, opts stt.Options) (string, error) {
|
||
args := m.Called(ctx, audio, opts)
|
||
return args.String(0), args.Error(1)
|
||
}
|
||
|
||
// MockLLMService mock LLM 服务
|
||
type MockLLMService struct {
|
||
mock.Mock
|
||
}
|
||
|
||
func (m *MockLLMService) ChatStream(ctx context.Context, req llm.Request) (<-chan llm.Chunk, error) {
|
||
args := m.Called(ctx, req)
|
||
if args.Get(0) == nil {
|
||
return nil, args.Error(1)
|
||
}
|
||
return args.Get(0).(<-chan llm.Chunk), args.Error(1)
|
||
}
|
||
|
||
// MockTTSService mock TTS 服务
|
||
type MockTTSService struct {
|
||
mock.Mock
|
||
}
|
||
|
||
func (m *MockTTSService) SynthesizeStream(ctx context.Context, textStream <-chan string, opts tts.Options) (<-chan tts.Chunk, error) {
|
||
args := m.Called(ctx, textStream, opts)
|
||
if args.Get(0) == nil {
|
||
return nil, args.Error(1)
|
||
}
|
||
return args.Get(0).(<-chan tts.Chunk), args.Error(1)
|
||
}
|
||
|
||
// MockSessionManager mock 会话管理器
|
||
type MockSessionManager struct {
|
||
mock.Mock
|
||
}
|
||
|
||
func (m *MockSessionManager) Create(ctx context.Context, config models.SessionConfig) (string, error) {
|
||
args := m.Called(ctx, config)
|
||
return args.String(0), args.Error(1)
|
||
}
|
||
|
||
func (m *MockSessionManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||
args := m.Called(ctx, sessionID)
|
||
if args.Get(0) == nil {
|
||
return nil, args.Error(1)
|
||
}
|
||
return args.Get(0).(*models.Session), args.Error(1)
|
||
}
|
||
|
||
func (m *MockSessionManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
||
args := m.Called(ctx, sessionID, patch)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||
args := m.Called(ctx, sessionID, limit)
|
||
return args.Get(0).([]models.Message), args.Error(1)
|
||
}
|
||
|
||
func (m *MockSessionManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
|
||
args := m.Called(ctx, sessionID, msg)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
|
||
args := m.Called(ctx, sessionID, requestID)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
|
||
args := m.Called(ctx, sessionID)
|
||
return args.String(0), args.Error(1)
|
||
}
|
||
|
||
func (m *MockSessionManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
|
||
args := m.Called(ctx, sessionID)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) Touch(ctx context.Context, sessionID string) error {
|
||
args := m.Called(ctx, sessionID)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) Destroy(ctx context.Context, sessionID string) error {
|
||
args := m.Called(ctx, sessionID)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSessionManager) ActiveCount() int {
|
||
args := m.Called()
|
||
return args.Int(0)
|
||
}
|
||
|
||
// MockSender mock WebSocket 发送器
|
||
type MockSender struct {
|
||
mock.Mock
|
||
STTResults []models.WsSTTResult
|
||
LLMChunks []models.WsLLMChunk
|
||
LLMDones []models.WsLLMDone
|
||
TTSAudios []models.WsTTSAudio
|
||
Errors []models.WsError
|
||
}
|
||
|
||
func NewMockSender() *MockSender {
|
||
return &MockSender{
|
||
STTResults: make([]models.WsSTTResult, 0),
|
||
LLMChunks: make([]models.WsLLMChunk, 0),
|
||
LLMDones: make([]models.WsLLMDone, 0),
|
||
TTSAudios: make([]models.WsTTSAudio, 0),
|
||
Errors: make([]models.WsError, 0),
|
||
}
|
||
}
|
||
|
||
func (m *MockSender) SendSTTResult(result models.WsSTTResult) error {
|
||
m.STTResults = append(m.STTResults, result)
|
||
args := m.Called(result)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSender) SendLLMChunk(chunk models.WsLLMChunk) error {
|
||
m.LLMChunks = append(m.LLMChunks, chunk)
|
||
args := m.Called(chunk)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSender) SendLLMDone(done models.WsLLMDone) error {
|
||
m.LLMDones = append(m.LLMDones, done)
|
||
args := m.Called(done)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSender) SendTTSAudio(audio models.WsTTSAudio) error {
|
||
m.TTSAudios = append(m.TTSAudios, audio)
|
||
args := m.Called(audio)
|
||
return args.Error(0)
|
||
}
|
||
|
||
func (m *MockSender) SendError(err models.WsError) error {
|
||
m.Errors = append(m.Errors, err)
|
||
args := m.Called(err)
|
||
return args.Error(0)
|
||
}
|
||
|
||
// 辅助函数:创建 LLM 流式响应
|
||
func createLLMStream(chunks []llm.Chunk) <-chan llm.Chunk {
|
||
ch := make(chan llm.Chunk, len(chunks))
|
||
for _, chunk := range chunks {
|
||
ch <- chunk
|
||
}
|
||
close(ch)
|
||
return ch
|
||
}
|
||
|
||
// 辅助函数:创建 TTS 流式响应
|
||
func createTTSStream(chunks []tts.Chunk) <-chan tts.Chunk {
|
||
ch := make(chan tts.Chunk, len(chunks))
|
||
for _, chunk := range chunks {
|
||
ch <- chunk
|
||
}
|
||
close(ch)
|
||
return ch
|
||
}
|
||
|
||
// TestProcessQuery_Success 测试完整流程
|
||
func TestProcessQuery_Success(t *testing.T) {
|
||
// 准备测试数据
|
||
audioData := []byte("test audio")
|
||
imageData := []byte("test image")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
imageBase64 := base64.StdEncoding.EncodeToString(imageData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Image: imageBase64,
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
session := &models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
TTSEnabled: true,
|
||
DetailLevel: "low",
|
||
Language: "zh-CN",
|
||
},
|
||
}
|
||
|
||
// 创建 mock
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
// 设置 mock 期望
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, stt.Options{
|
||
Encoding: "pcm_s16le",
|
||
SampleRate: 16000,
|
||
Language: "zh-CN",
|
||
}).Return("你好,世界", nil)
|
||
|
||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
||
|
||
llmChunks := []llm.Chunk{
|
||
{Delta: "你好"},
|
||
{Delta: ",世界!"},
|
||
{Done: true, TokensUsed: &llm.TokenUsage{Prompt: 10, Completion: 5, Total: 15}},
|
||
}
|
||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
||
|
||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
||
|
||
ttsChunks := []tts.Chunk{
|
||
{Audio: []byte("audio1"), IsLast: false},
|
||
{Audio: []byte("audio2"), IsLast: true},
|
||
}
|
||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).Return(createTTSStream(ttsChunks), nil)
|
||
|
||
mockSender.On("SendTTSAudio", mock.Anything).Return(nil)
|
||
|
||
// 创建 Pipeline
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
// 执行
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
// 验证
|
||
assert.NoError(t, err)
|
||
assert.Len(t, mockSender.STTResults, 1)
|
||
assert.Equal(t, "你好,世界", mockSender.STTResults[0].Text)
|
||
assert.Len(t, mockSender.LLMChunks, 2)
|
||
assert.Len(t, mockSender.LLMDones, 1)
|
||
assert.Len(t, mockSender.TTSAudios, 2)
|
||
|
||
mockSTT.AssertExpectations(t)
|
||
mockLLM.AssertExpectations(t)
|
||
mockTTS.AssertExpectations(t)
|
||
mockSession.AssertExpectations(t)
|
||
}
|
||
|
||
// TestProcessQuery_STTError 测试 STT 失败降级
|
||
func TestProcessQuery_STTError(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(&models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
Language: "zh-CN",
|
||
},
|
||
}, nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).
|
||
Return("", errors.New("STT service unavailable"))
|
||
|
||
mockSender.On("SendError", mock.Anything).Return(nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
assert.Error(t, err)
|
||
assert.Len(t, mockSender.Errors, 1)
|
||
assert.Equal(t, "STT_ERROR", mockSender.Errors[0].Code)
|
||
|
||
mockSTT.AssertExpectations(t)
|
||
mockLLM.AssertNotCalled(t, "ChatStream")
|
||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
||
}
|
||
|
||
// TestProcessQuery_LLMError 测试 LLM 失败降级
|
||
func TestProcessQuery_LLMError(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
session := &models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
TTSEnabled: true,
|
||
Language: "zh-CN",
|
||
},
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
||
|
||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).
|
||
Return(nil, errors.New("LLM service unavailable"))
|
||
|
||
mockSender.On("SendError", mock.Anything).Return(nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
assert.Error(t, err)
|
||
assert.Len(t, mockSender.Errors, 1)
|
||
assert.Equal(t, "LLM_ERROR", mockSender.Errors[0].Code)
|
||
|
||
mockSTT.AssertExpectations(t)
|
||
mockLLM.AssertExpectations(t)
|
||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
||
}
|
||
|
||
// TestProcessQuery_TTSError 测试 TTS 失败静默跳过
|
||
func TestProcessQuery_TTSError(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
session := &models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
TTSEnabled: true,
|
||
Language: "zh-CN",
|
||
},
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
||
|
||
llmChunks := []llm.Chunk{
|
||
{Delta: "你好"},
|
||
{Done: true},
|
||
}
|
||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
||
|
||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).
|
||
Return(nil, errors.New("TTS service unavailable"))
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
// TTS 失败应该静默跳过,不返回错误
|
||
assert.NoError(t, err)
|
||
assert.Len(t, mockSender.LLMDones, 1)
|
||
assert.Len(t, mockSender.TTSAudios, 0)
|
||
|
||
mockSTT.AssertExpectations(t)
|
||
mockLLM.AssertExpectations(t)
|
||
mockTTS.AssertExpectations(t)
|
||
}
|
||
|
||
// TestProcessQuery_ContextCancelled 测试上下文取消(Interrupt)
|
||
func TestProcessQuery_ContextCancelled(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
session := &models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
TTSEnabled: true,
|
||
Language: "zh-CN",
|
||
},
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
||
|
||
// 创建一个会延迟的 LLM 流,以便我们可以取消上下文
|
||
llmCh := make(chan llm.Chunk)
|
||
go func() {
|
||
time.Sleep(100 * time.Millisecond)
|
||
llmCh <- llm.Chunk{Delta: "你"}
|
||
time.Sleep(100 * time.Millisecond)
|
||
llmCh <- llm.Chunk{Delta: "好"}
|
||
close(llmCh)
|
||
}()
|
||
|
||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return((<-chan llm.Chunk)(llmCh), nil)
|
||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
||
|
||
// 创建一个会延迟的 TTS 流
|
||
ttsCh := make(chan tts.Chunk)
|
||
go func() {
|
||
time.Sleep(200 * time.Millisecond)
|
||
close(ttsCh)
|
||
}()
|
||
mockTTS.On("SynthesizeStream", mock.Anything, mock.Anything, mock.Anything).Return((<-chan tts.Chunk)(ttsCh), nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
// 创建可取消的上下文
|
||
ctx, cancel := context.WithCancel(context.Background())
|
||
|
||
// 在 50ms 后取消
|
||
go func() {
|
||
time.Sleep(50 * time.Millisecond)
|
||
cancel()
|
||
}()
|
||
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
// 上下文取消后,流程应该正常完成(中断流但不返回错误)
|
||
assert.NoError(t, err)
|
||
|
||
mockSTT.AssertExpectations(t)
|
||
}
|
||
|
||
// TestProcessQuery_DisabledTTS 测试 TTS 未启用的情况
|
||
func TestProcessQuery_DisabledTTS(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
session := &models.Session{
|
||
ID: "session-123",
|
||
Config: models.SessionConfig{
|
||
TTSEnabled: false, // TTS 未启用
|
||
Language: "zh-CN",
|
||
},
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(session, nil)
|
||
mockSession.On("AppendMessage", mock.Anything, "session-123", mock.Anything).Return(nil)
|
||
|
||
mockSTT.On("Recognize", mock.Anything, audioData, mock.Anything).Return("你好", nil)
|
||
mockSender.On("SendSTTResult", mock.Anything).Return(nil)
|
||
|
||
llmChunks := []llm.Chunk{
|
||
{Delta: "你好"},
|
||
{Done: true},
|
||
}
|
||
mockLLM.On("ChatStream", mock.Anything, mock.Anything).Return(createLLMStream(llmChunks), nil)
|
||
mockSender.On("SendLLMChunk", mock.Anything).Return(nil)
|
||
mockSender.On("SendLLMDone", mock.Anything).Return(nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
assert.NoError(t, err)
|
||
assert.Len(t, mockSender.LLMDones, 1)
|
||
assert.Len(t, mockSender.TTSAudios, 0)
|
||
|
||
// TTS 不应该被调用
|
||
mockTTS.AssertNotCalled(t, "SynthesizeStream")
|
||
}
|
||
|
||
// TestSplitter 测试句子切分器
|
||
func TestSplitter(t *testing.T) {
|
||
ch := make(chan string, 10)
|
||
splitter := NewSplitter(ch)
|
||
|
||
// 输入包含多个句子的文本
|
||
splitter.Feed("你好。")
|
||
splitter.Feed("世界!")
|
||
splitter.Feed("这是")
|
||
splitter.Feed("一个测试。")
|
||
splitter.Flush()
|
||
|
||
// 应该有 3 个句子
|
||
assert.Equal(t, 3, len(ch))
|
||
assert.Equal(t, "你好。", <-ch)
|
||
assert.Equal(t, "世界!", <-ch)
|
||
assert.Equal(t, "这是一个测试。", <-ch)
|
||
}
|
||
|
||
// TestSplitter_NoDelimiter 测试没有分隔符的情况
|
||
func TestSplitter_NoDelimiter(t *testing.T) {
|
||
ch := make(chan string, 10)
|
||
splitter := NewSplitter(ch)
|
||
|
||
splitter.Feed("没有分隔符的文本")
|
||
splitter.Flush()
|
||
|
||
// 应该有 1 个句子(Flush 会发送剩余内容)
|
||
assert.Equal(t, 1, len(ch))
|
||
assert.Equal(t, "没有分隔符的文本", <-ch)
|
||
}
|
||
|
||
// TestSplitter_Empty 测试空输入
|
||
func TestSplitter_Empty(t *testing.T) {
|
||
ch := make(chan string, 10)
|
||
splitter := NewSplitter(ch)
|
||
|
||
splitter.Flush()
|
||
|
||
// 应该没有句子
|
||
assert.Equal(t, 0, len(ch))
|
||
}
|
||
|
||
// TestProcessQuery_InvalidAudio 测试无效音频数据
|
||
func TestProcessQuery_InvalidAudio(t *testing.T) {
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: "invalid-base64!!!",
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSender.On("SendError", mock.Anything).Return(nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
assert.Error(t, err)
|
||
assert.Len(t, mockSender.Errors, 1)
|
||
assert.Equal(t, "INVALID_MESSAGE", mockSender.Errors[0].Code)
|
||
}
|
||
|
||
// TestProcessQuery_SessionNotFound 测试会话不存在
|
||
func TestProcessQuery_SessionNotFound(t *testing.T) {
|
||
audioData := []byte("test audio")
|
||
audioBase64 := base64.StdEncoding.EncodeToString(audioData)
|
||
|
||
req := models.WsQuery{
|
||
Type: "query",
|
||
RequestID: "req-123",
|
||
Audio: audioBase64,
|
||
}
|
||
|
||
mockSTT := new(MockSTTService)
|
||
mockLLM := new(MockLLMService)
|
||
mockTTS := new(MockTTSService)
|
||
mockSession := new(MockSessionManager)
|
||
mockSender := NewMockSender()
|
||
|
||
mockSession.On("SetActiveRequest", mock.Anything, "session-123", "req-123").Return(nil)
|
||
mockSession.On("ClearActiveRequest", mock.Anything, "session-123").Return(nil)
|
||
mockSession.On("Get", mock.Anything, "session-123").Return(nil, errors.New("session not found"))
|
||
|
||
mockSender.On("SendError", mock.Anything).Return(nil)
|
||
|
||
pipeline := New(mockSTT, mockLLM, mockTTS, mockSession, "gpt-4o", "alloy", 1.0)
|
||
|
||
ctx := context.Background()
|
||
err := pipeline.ProcessQuery(ctx, "session-123", req, nil, mockSender)
|
||
|
||
assert.Error(t, err)
|
||
assert.Len(t, mockSender.Errors, 1)
|
||
assert.Equal(t, "SESSION_NOT_FOUND", mockSender.Errors[0].Code)
|
||
}
|