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 }