Files
fire-safety-ymd/internal/service/chat_test.go
2026-09-05 15:46:37 +08:00

279 lines
8.8 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"strings"
"sync"
"testing"
"time"
)
func TestChatServiceCreatesAndReusesProviderSession(t *testing.T) {
agent := &fakeChatAgent{}
chat := newTestChatService(t, agent, nil, 10)
first, err := chat.Prepare(context.Background(), ChatRequest{Message: " 第一问 ", RequestID: "request-1"})
if err != nil {
t.Fatalf("Prepare(first) error = %v", err)
}
if first.Reused() || !conversationIDPattern.MatchString(first.ConversationID()) {
t.Fatalf("first turn = id %q reused %t", first.ConversationID(), first.Reused())
}
firstResult, err := first.Stream(context.Background(), func(event AgentTraceEvent) {
if event.Event != "tool.started" || event.ToolName != "fire_safety_test" {
t.Errorf("trace event = %#v", event)
}
})
if err != nil {
t.Fatalf("Stream(first) error = %v", err)
}
if firstResult.Answer != "answer-1" || firstResult.ConversationID != first.ConversationID() {
t.Fatalf("first result = %#v", firstResult)
}
if firstResult.ModelID != "test-model" {
t.Fatalf("first model ID = %q", firstResult.ModelID)
}
second, err := chat.Prepare(context.Background(), ChatRequest{
Message: "第二问",
ConversationID: first.ConversationID(),
RequestID: "request-2",
})
if err != nil {
t.Fatalf("Prepare(second) error = %v", err)
}
if !second.Reused() {
t.Fatal("second.Reused() = false")
}
if _, err := second.Stream(context.Background(), nil); err != nil {
t.Fatalf("Stream(second) error = %v", err)
}
agent.mu.Lock()
defer agent.mu.Unlock()
if len(agent.creates) != 1 || len(agent.messages) != 2 {
t.Fatalf("provider calls: creates=%d messages=%d", len(agent.creates), len(agent.messages))
}
if agent.creates[0].ExternalSubjectID != "chat-test-subject" {
t.Fatalf("ExternalSubjectID = %q", agent.creates[0].ExternalSubjectID)
}
if agent.messages[0].SessionID != agent.messages[1].SessionID {
t.Fatalf("provider sessions differ: %q and %q", agent.messages[0].SessionID, agent.messages[1].SessionID)
}
if agent.messages[0].Message != "第一问" || agent.messages[1].Message != "第二问" {
t.Fatalf("provider messages = %q, %q", agent.messages[0].Message, agent.messages[1].Message)
}
if agent.messages[0].IdempotencyKey == agent.messages[1].IdempotencyKey {
t.Fatal("turn idempotency keys were reused")
}
}
func TestChatServiceRejectsConcurrentRunOnSameConversation(t *testing.T) {
block := make(chan struct{})
started := make(chan struct{})
agent := &fakeChatAgent{block: block, started: started}
chat := newTestChatService(t, agent, nil, 10)
turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"})
if err != nil {
t.Fatalf("Prepare() error = %v", err)
}
done := make(chan error, 1)
go func() {
_, streamErr := turn.Stream(context.Background(), nil)
done <- streamErr
}()
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("provider stream did not start")
}
_, err = chat.Prepare(context.Background(), ChatRequest{
Message: "overlap",
ConversationID: turn.ConversationID(),
RequestID: "request-2",
})
if !errors.Is(err, ErrChatConversationBusy) {
t.Fatalf("Prepare(overlap) error = %v, want busy", err)
}
close(block)
if err := <-done; err != nil {
t.Fatalf("Stream() error = %v", err)
}
next, err := chat.Prepare(context.Background(), ChatRequest{
Message: "after completion",
ConversationID: turn.ConversationID(),
RequestID: "request-3",
})
if err != nil {
t.Fatalf("Prepare(after completion) error = %v", err)
}
next.Close()
}
func TestChatServiceInvalidatesConversationAfterUncertainStreamFailure(t *testing.T) {
agent := &fakeChatAgent{streamErr: ErrChatUpstreamProtocol}
chat := newTestChatService(t, agent, nil, 10)
turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"})
if err != nil {
t.Fatalf("Prepare() error = %v", err)
}
if _, err := turn.Stream(context.Background(), nil); !errors.Is(err, ErrChatUpstreamProtocol) {
t.Fatalf("Stream() error = %v", err)
}
_, err = chat.Prepare(context.Background(), ChatRequest{
Message: "retry",
ConversationID: turn.ConversationID(),
RequestID: "request-2",
})
if !errors.Is(err, ErrChatConversationNotFound) {
t.Fatalf("Prepare(retry) error = %v, want not found", err)
}
}
func TestChatServiceExpiresIdleSessionsAndEnforcesCapacity(t *testing.T) {
now := time.Date(2026, 9, 5, 12, 0, 0, 0, time.UTC)
clock := func() time.Time { return now }
agent := &fakeChatAgent{}
chat := newTestChatService(t, agent, clock, 1)
first, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"})
if err != nil {
t.Fatalf("Prepare(first) error = %v", err)
}
if _, err := first.Stream(context.Background(), nil); err != nil {
t.Fatalf("Stream(first) error = %v", err)
}
if _, err := chat.Prepare(context.Background(), ChatRequest{Message: "new", RequestID: "request-2"}); !errors.Is(err, ErrChatCapacityReached) {
t.Fatalf("Prepare(at capacity) error = %v", err)
}
now = now.Add(31 * time.Minute)
second, err := chat.Prepare(context.Background(), ChatRequest{Message: "new", RequestID: "request-3"})
if err != nil {
t.Fatalf("Prepare(after TTL) error = %v", err)
}
second.Close()
_, err = chat.Prepare(context.Background(), ChatRequest{
Message: "old",
ConversationID: first.ConversationID(),
RequestID: "request-4",
})
if !errors.Is(err, ErrChatConversationNotFound) {
t.Fatalf("Prepare(expired) error = %v", err)
}
}
func TestChatServiceReleasesReservationAfterCreateFailure(t *testing.T) {
agent := &fakeChatAgent{createErr: ErrChatUpstreamUnavailable}
chat := newTestChatService(t, agent, nil, 1)
if _, err := chat.Prepare(context.Background(), ChatRequest{Message: "first", RequestID: "request-1"}); !errors.Is(err, ErrChatUpstreamUnavailable) {
t.Fatalf("Prepare(first) error = %v", err)
}
agent.mu.Lock()
agent.createErr = nil
agent.mu.Unlock()
turn, err := chat.Prepare(context.Background(), ChatRequest{Message: "second", RequestID: "request-2"})
if err != nil {
t.Fatalf("Prepare(second) error = %v", err)
}
turn.Close()
}
func TestChatServiceRejectsInvalidRequests(t *testing.T) {
chat := newTestChatService(t, &fakeChatAgent{}, nil, 10)
tests := []ChatRequest{
{Message: "", RequestID: "request-1"},
{Message: strings.Repeat("a", 1025), RequestID: "request-1"},
{Message: "ok", RequestID: "invalid request id"},
{Message: "ok", RequestID: "request-1", ConversationID: "provider-session-1"},
}
for _, request := range tests {
if _, err := chat.Prepare(context.Background(), request); !errors.Is(err, ErrChatInvalidArgument) {
t.Fatalf("Prepare(%#v) error = %v, want invalid argument", request, err)
}
}
}
func newTestChatService(t *testing.T, agent ChatAgent, now func() time.Time, maxSessions int) *ChatService {
t.Helper()
var idMu sync.Mutex
idSequence := 0
newID := func(prefix string) (string, error) {
idMu.Lock()
defer idMu.Unlock()
idSequence++
return fmt.Sprintf("%s%024d", prefix, idSequence), nil
}
chat, err := NewChatService(agent, ChatOptions{
ExternalSubjectID: "chat-test-subject",
MaxMessageBytes: 1024,
SessionTTL: 30 * time.Minute,
MaxSessions: maxSessions,
now: now,
newID: newID,
})
if err != nil {
t.Fatalf("NewChatService() error = %v", err)
}
return chat
}
type fakeChatAgent struct {
mu sync.Mutex
creates []AgentCreateSessionRequest
messages []AgentMessageRequest
createErr error
streamErr error
block <-chan struct{}
started chan<- struct{}
startOnce sync.Once
sessionSeq int
}
func (a *fakeChatAgent) CreateSession(_ context.Context, request AgentCreateSessionRequest) (AgentSession, error) {
a.mu.Lock()
defer a.mu.Unlock()
a.creates = append(a.creates, request)
if a.createErr != nil {
return AgentSession{}, a.createErr
}
a.sessionSeq++
return AgentSession{ID: fmt.Sprintf("provider-session-%d", a.sessionSeq)}, nil
}
func (a *fakeChatAgent) StreamMessage(ctx context.Context, request AgentMessageRequest, trace func(AgentTraceEvent)) (AgentMessageResult, error) {
a.mu.Lock()
a.messages = append(a.messages, request)
messageNumber := len(a.messages)
streamErr := a.streamErr
block := a.block
started := a.started
a.mu.Unlock()
if started != nil {
a.startOnce.Do(func() { close(started) })
}
if block != nil {
select {
case <-ctx.Done():
return AgentMessageResult{}, ctx.Err()
case <-block:
}
}
if trace != nil {
trace(AgentTraceEvent{Event: "tool.started", ToolName: "fire_safety_test", Status: "running"})
}
if streamErr != nil {
return AgentMessageResult{}, streamErr
}
return AgentMessageResult{
RunID: fmt.Sprintf("run-%d", messageNumber),
ModelID: "test-model",
Answer: fmt.Sprintf("answer-%d", messageNumber),
Usage: ChatTokenUsage{Input: 1, Output: 2, Total: 3},
}, nil
}