package eino import ( "context" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "github.com/hhs/camtalk/internal/ai/stt" "github.com/hhs/camtalk/internal/ai/tts" "github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/orchestrator" "github.com/hhs/camtalk/internal/trace" ) // --- Mock STT Service --- 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) } // --- Mock TTS Service --- 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) return args.Get(0).(<-chan tts.Chunk), args.Error(1) } // --- Mock Sender --- type mockSender struct { mock.Mock STTResults []models.WsSTTResult LLMChunks []models.WsLLMChunk LLMDones []models.WsLLMDone TTSAudios []models.WsTTSAudio Errors []models.WsError } func (m *mockSender) SendSTTResult(result models.WsSTTResult) error { m.STTResults = append(m.STTResults, result) return m.Called(result).Error(0) } func (m *mockSender) SendLLMChunk(chunk models.WsLLMChunk) error { m.LLMChunks = append(m.LLMChunks, chunk) return m.Called(chunk).Error(0) } func (m *mockSender) SendLLMDone(done models.WsLLMDone) error { m.LLMDones = append(m.LLMDones, done) return m.Called(done).Error(0) } func (m *mockSender) SendTTSAudio(audio models.WsTTSAudio) error { m.TTSAudios = append(m.TTSAudios, audio) return m.Called(audio).Error(0) } func (m *mockSender) SendError(err models.WsError) error { m.Errors = append(m.Errors, err) return m.Called(err).Error(0) } // --- Tests --- func TestDetectImageMimeType(t *testing.T) { tests := []struct { name string data []byte expected string }{ {"JPEG", []byte{0xFF, 0xD8, 0xFF, 0xE0}, "image/jpeg"}, {"PNG", []byte{0x89, 0x50, 0x4E, 0x47}, "image/png"}, {"GIF", []byte{0x47, 0x49, 0x46, 0x38}, "image/gif"}, {"WebP", []byte{0x52, 0x49, 0x46, 0x46}, "image/webp"}, {"Unknown", []byte{0x00, 0x00, 0x00}, "image/jpeg"}, {"Short", []byte{0xFF}, "image/jpeg"}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { result := detectImageMimeType(tt.data) assert.Equal(t, tt.expected, result) }) } } func TestBuildPipelineInput(t *testing.T) { req := models.WsQuery{ Text: "你好", RequestID: "req-1", } sess := &models.Session{ Config: models.SessionConfig{ Language: "zh-CN", Scenario: "free_chat", TTSEnabled: true, }, } input := buildPipelineInput(req, "sess-1", sess, nil, nil) require.Equal(t, "你好", input.Text) require.Equal(t, "sess-1", input.SessionID) require.Equal(t, "req-1", input.RequestID) require.Equal(t, "zh-CN", input.Language) require.Equal(t, "free_chat", input.Scenario) require.True(t, input.TTSEnabled) } func TestBuildPipelineInput_WithAudioData(t *testing.T) { req := models.WsQuery{ Audio: "base64audio", RequestID: "req-2", } sess := &models.Session{ Config: models.SessionConfig{ Language: "en", Scenario: "free_chat", TTSEnabled: false, }, } audioData := []byte("fake-audio-bytes") imageData := []byte("fake-image-bytes") input := buildPipelineInput(req, "sess-2", sess, audioData, imageData) require.Equal(t, audioData, input.AudioData) require.Equal(t, imageData, input.ImageData) require.False(t, input.TTSEnabled) require.Equal(t, "en", input.Language) } func TestPipelineState_AppendAndGet(t *testing.T) { state := genLocalState(context.Background()) state.AppendText("Hello ") state.AppendText("World") require.Equal(t, "Hello World", state.GetFullResponse()) } func TestPipelineState_ConcurrentAccess(t *testing.T) { state := genLocalState(context.Background()) done := make(chan struct{}) go func() { for i := 0; i < 100; i++ { state.AppendText("a") } close(done) }() for i := 0; i < 100; i++ { _ = state.GetFullResponse() } <-done require.Equal(t, 100, len(state.GetFullResponse())) } func TestContextInjection(t *testing.T) { ctx := context.Background() sender := &mockSender{} ctx = WithSender(ctx, sender) ctx = WithRequestID(ctx, "req-123") ctx = trace.WithSessionID(ctx, "sess-456") ctx = WithStartTime(ctx, time.Now()) ctx = WithPipelineState(ctx, genLocalState(ctx)) require.NotNil(t, senderFromCtx(ctx)) require.Equal(t, "req-123", requestIDFromCtx(ctx)) require.NotNil(t, stateFromCtx(ctx)) } func TestLatencyFromCtx(t *testing.T) { ctx := context.Background() // No start time set require.Equal(t, int64(0), latencyFromCtx(ctx)) // With start time start := time.Now().Add(-100 * time.Millisecond) ctx = WithStartTime(ctx, start) latency := latencyFromCtx(ctx) require.Greater(t, latency, int64(0)) require.Less(t, latency, int64(1000)) // should be < 1 second } func TestEinoOrchestrator_ImplementsInterface(t *testing.T) { // Compile-time check that EinoOrchestrator implements orchestrator.Orchestrator var _ orchestrator.Orchestrator = (*EinoOrchestrator)(nil) } func TestNewSTTLambda_ReturnsNonNil(t *testing.T) { mockSTT := &mockSTTService{} lambda := NewSTTLambda(mockSTT) require.NotNil(t, lambda) } func TestNewHistoryLambda_ReturnsNonNil(t *testing.T) { fetcher := func(ctx context.Context, sessionID string, limit int) ([]models.Message, error) { return nil, nil } lambda := NewHistoryLambda(fetcher, nil, 10) require.NotNil(t, lambda) } func TestNewSplitterLambda_ReturnsNonNil(t *testing.T) { lambda := NewSplitterLambda() require.NotNil(t, lambda) } func TestNewTTSLambda_ReturnsNonNil(t *testing.T) { mockTTS := &mockTTSService{} lambda := NewTTSLambda(mockTTS, "alloy", 1.0, "mp3", 24000) require.NotNil(t, lambda) } func TestNewDoneLambda_ReturnsNonNil(t *testing.T) { lambda := NewDoneLambda("test-model") require.NotNil(t, lambda) }