Files

237 lines
5.9 KiB
Go
Raw Permalink Normal View History

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)
}