279 lines
8.8 KiB
Go
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
|
|
}
|