package service import ( "context" "crypto/rand" "encoding/base64" "errors" "fmt" "regexp" "strings" "sync" "time" "unicode" ) const ( maximumChatMessageBytes int64 = 16 * 1024 * 1024 maximumChatSessionTTL = 24 * time.Hour maximumChatSessions = 10_000 ) var ( // ErrChatInvalidArgument indicates an invalid user-facing chat request. ErrChatInvalidArgument = errors.New("invalid chat argument") // ErrChatConversationNotFound indicates an unknown, expired, or lost in-memory conversation. ErrChatConversationNotFound = errors.New("chat conversation not found") // ErrChatConversationBusy indicates that a conversation already has an active run. ErrChatConversationBusy = errors.New("chat conversation is busy") // ErrChatCapacityReached indicates that the bounded in-memory conversation store is full. ErrChatCapacityReached = errors.New("chat conversation capacity reached") // ErrChatUpstreamUnavailable indicates a provider transport or availability failure. ErrChatUpstreamUnavailable = errors.New("chat upstream unavailable") // ErrChatUpstreamProtocol indicates that the provider returned an unsafe or incomplete protocol result. ErrChatUpstreamProtocol = errors.New("chat upstream protocol error") // ErrChatRunFailed indicates that the provider explicitly reported a failed run. ErrChatRunFailed = errors.New("chat upstream run failed") // ErrChatInternal indicates a local failure that cannot be attributed to user input. ErrChatInternal = errors.New("chat internal failure") conversationIDPattern = regexp.MustCompile(`^conv_[A-Za-z0-9_-]{16,128}$`) ) // ChatAgent is the provider-neutral outbound boundary used by ChatService. type ChatAgent interface { CreateSession(context.Context, AgentCreateSessionRequest) (AgentSession, error) StreamMessage(context.Context, AgentMessageRequest, func(AgentTraceEvent)) (AgentMessageResult, error) } // AgentCreateSessionRequest contains only server-controlled identity and metadata. type AgentCreateSessionRequest struct { ExternalSubjectID string IdempotencyKey string RequestID string Metadata map[string]any } // AgentSession is the provider's opaque session handle. type AgentSession struct { ID string } // AgentMessageRequest sends one turn to an already-created provider session. type AgentMessageRequest struct { SessionID string Message string IdempotencyKey string RequestID string Metadata map[string]any } // AgentTraceEvent is the deliberately small public progress projection. It // excludes text, identifiers, tool inputs, tool outputs, and provider payloads. type AgentTraceEvent struct { Event string `json:"event"` ToolName string `json:"tool_name,omitempty"` Status string `json:"status,omitempty"` } // ChatTokenUsage contains non-negative provider-reported token counts. type ChatTokenUsage struct { Input int64 `json:"input"` Output int64 `json:"output"` Total int64 `json:"total"` } // AgentMessageResult is returned by a ChatAgent only after strict provider success. type AgentMessageResult struct { RunID string ModelID string Answer string Usage ChatTokenUsage } // ChatRequest is the validated application-level request for one turn. type ChatRequest struct { Message string ConversationID string RequestID string } // ChatResult is safe for the inbound handler to expose after strict success. type ChatResult struct { ConversationID string `json:"conversation_id"` RunID string `json:"run_id,omitempty"` ModelID string `json:"model_id,omitempty"` Answer string `json:"answer"` Usage ChatTokenUsage `json:"usage"` } // ChatTurn represents an exclusively acquired conversation turn. Call Close if // Stream is not called; Stream releases or invalidates the conversation itself. type ChatTurn interface { ConversationID() string Reused() bool Stream(context.Context, func(AgentTraceEvent)) (ChatResult, error) Close() } // ChatOptions controls the bounded, process-local conversation store. type ChatOptions struct { ExternalSubjectID string MaxMessageBytes int64 SessionTTL time.Duration MaxSessions int now func() time.Time newID func(string) (string, error) } // ChatService owns local-to-provider session mapping and per-conversation run exclusion. type ChatService struct { agent ChatAgent externalSubjectID string maxMessageBytes int64 sessionTTL time.Duration maxSessions int now func() time.Time newID func(string) (string, error) mu sync.Mutex conversations map[string]*chatConversation reserved map[string]struct{} } type chatConversation struct { providerSessionID string busy bool lastUsed time.Time } // NewChatService constructs a bounded in-memory chat service. func NewChatService(agent ChatAgent, options ChatOptions) (*ChatService, error) { if agent == nil { return nil, errors.New("chat agent is required") } options.ExternalSubjectID = strings.TrimSpace(options.ExternalSubjectID) if options.ExternalSubjectID == "" || len(options.ExternalSubjectID) > 512 || containsControl(options.ExternalSubjectID) { return nil, errors.New("chat external subject ID is invalid") } if options.MaxMessageBytes <= 0 || options.MaxMessageBytes > maximumChatMessageBytes { return nil, errors.New("chat message byte limit is outside the supported range") } if options.SessionTTL <= 0 || options.SessionTTL > maximumChatSessionTTL { return nil, errors.New("chat session TTL is outside the supported range") } if options.MaxSessions <= 0 || options.MaxSessions > maximumChatSessions { return nil, errors.New("chat session capacity is outside the supported range") } if options.now == nil { options.now = time.Now } if options.newID == nil { options.newID = randomChatID } return &ChatService{ agent: agent, externalSubjectID: options.ExternalSubjectID, maxMessageBytes: options.MaxMessageBytes, sessionTTL: options.SessionTTL, maxSessions: options.MaxSessions, now: options.now, newID: options.newID, conversations: make(map[string]*chatConversation), reserved: make(map[string]struct{}), }, nil } // Prepare validates a request and exclusively acquires either a new or existing conversation. func (s *ChatService) Prepare(ctx context.Context, request ChatRequest) (ChatTurn, error) { message := strings.TrimSpace(request.Message) if message == "" || int64(len([]byte(message))) > s.maxMessageBytes { return nil, fmt.Errorf("%w: message is required and bounded", ErrChatInvalidArgument) } requestID := strings.TrimSpace(request.RequestID) if requestID == "" { generated, err := s.newID("req_") if err != nil { return nil, fmt.Errorf("%w: generate request ID", ErrChatInternal) } requestID = generated } if !validCorrelationID(requestID) { return nil, fmt.Errorf("%w: request ID is invalid", ErrChatInvalidArgument) } conversationID := strings.TrimSpace(request.ConversationID) if conversationID != "" { if !conversationIDPattern.MatchString(conversationID) { return nil, fmt.Errorf("%w: conversation ID is invalid", ErrChatInvalidArgument) } return s.prepareExisting(message, conversationID, requestID) } return s.prepareNew(ctx, message, requestID) } func (s *ChatService) prepareExisting(message, conversationID, requestID string) (ChatTurn, error) { now := s.now() s.mu.Lock() s.evictExpiredLocked(now) record, exists := s.conversations[conversationID] if !exists { s.mu.Unlock() return nil, ErrChatConversationNotFound } if record.busy { s.mu.Unlock() return nil, ErrChatConversationBusy } record.busy = true record.lastUsed = now s.mu.Unlock() idempotencyKey, err := s.newID("turn_") if err != nil || !validCorrelationID(idempotencyKey) { s.release(conversationID, record, true) return nil, fmt.Errorf("%w: generate turn ID", ErrChatInternal) } return &chatTurn{ service: s, record: record, conversationID: conversationID, message: message, requestID: requestID, idempotencyKey: idempotencyKey, reused: true, }, nil } func (s *ChatService) prepareNew(ctx context.Context, message, requestID string) (ChatTurn, error) { conversationID, err := s.reserveConversationID() if err != nil { return nil, err } releaseReservation := true defer func() { if releaseReservation { s.mu.Lock() delete(s.reserved, conversationID) s.mu.Unlock() } }() createKey, err := s.newID("session_") if err != nil || !validCorrelationID(createKey) { return nil, fmt.Errorf("%w: generate session ID", ErrChatInternal) } session, err := s.agent.CreateSession(ctx, AgentCreateSessionRequest{ ExternalSubjectID: s.externalSubjectID, IdempotencyKey: createKey, RequestID: requestID, Metadata: map[string]any{"source": "fire-safety-ymd-chat-api", "api_version": "v1"}, }) if err != nil { return nil, err } if !validOpaqueProviderID(session.ID) { return nil, fmt.Errorf("%w: provider returned an invalid session", ErrChatUpstreamProtocol) } idempotencyKey, err := s.newID("turn_") if err != nil || !validCorrelationID(idempotencyKey) { return nil, fmt.Errorf("%w: generate turn ID", ErrChatInternal) } record := &chatConversation{providerSessionID: session.ID, busy: true, lastUsed: s.now()} s.mu.Lock() delete(s.reserved, conversationID) s.conversations[conversationID] = record s.mu.Unlock() releaseReservation = false return &chatTurn{ service: s, record: record, conversationID: conversationID, message: message, requestID: requestID, idempotencyKey: idempotencyKey, }, nil } func (s *ChatService) reserveConversationID() (string, error) { for attempt := 0; attempt < 4; attempt++ { conversationID, err := s.newID("conv_") if err != nil || !conversationIDPattern.MatchString(conversationID) { return "", fmt.Errorf("%w: generate conversation ID", ErrChatInternal) } now := s.now() s.mu.Lock() s.evictExpiredLocked(now) if len(s.conversations)+len(s.reserved) >= s.maxSessions { s.mu.Unlock() return "", ErrChatCapacityReached } _, conversationExists := s.conversations[conversationID] _, reservationExists := s.reserved[conversationID] if !conversationExists && !reservationExists { s.reserved[conversationID] = struct{}{} s.mu.Unlock() return conversationID, nil } s.mu.Unlock() } return "", fmt.Errorf("%w: conversation ID collision", ErrChatInternal) } func (s *ChatService) evictExpiredLocked(now time.Time) { for conversationID, record := range s.conversations { if !record.busy && now.Sub(record.lastUsed) >= s.sessionTTL { delete(s.conversations, conversationID) } } } func (s *ChatService) release(conversationID string, record *chatConversation, keep bool) { s.mu.Lock() defer s.mu.Unlock() current, exists := s.conversations[conversationID] if !exists || current != record { return } if !keep { delete(s.conversations, conversationID) return } record.busy = false record.lastUsed = s.now() } type chatTurn struct { service *ChatService record *chatConversation conversationID string message string requestID string idempotencyKey string reused bool streamMu sync.Mutex streamed bool closed bool finish sync.Once } func (t *chatTurn) ConversationID() string { return t.conversationID } func (t *chatTurn) Reused() bool { return t.reused } func (t *chatTurn) Stream(ctx context.Context, traceHandler func(AgentTraceEvent)) (ChatResult, error) { t.streamMu.Lock() if t.streamed || t.closed { t.streamMu.Unlock() return ChatResult{}, fmt.Errorf("%w: turn is no longer available", ErrChatInternal) } t.streamed = true t.streamMu.Unlock() result, err := t.service.agent.StreamMessage(ctx, AgentMessageRequest{ SessionID: t.record.providerSessionID, Message: t.message, IdempotencyKey: t.idempotencyKey, RequestID: t.requestID, Metadata: map[string]any{"source": "fire-safety-ymd-chat-api", "api_version": "v1"}, }, traceHandler) if err != nil { t.complete(false) return ChatResult{}, err } if strings.TrimSpace(result.Answer) == "" || result.Usage.Input < 0 || result.Usage.Output < 0 || result.Usage.Total < 0 { t.complete(false) return ChatResult{}, fmt.Errorf("%w: provider returned an invalid final result", ErrChatUpstreamProtocol) } t.complete(true) return ChatResult{ ConversationID: t.conversationID, RunID: result.RunID, ModelID: result.ModelID, Answer: result.Answer, Usage: result.Usage, }, nil } func (t *chatTurn) Close() { t.streamMu.Lock() if t.streamed { t.closed = true t.streamMu.Unlock() return } t.closed = true keep := t.reused t.streamMu.Unlock() t.complete(keep) } func (t *chatTurn) complete(keep bool) { t.finish.Do(func() { t.service.release(t.conversationID, t.record, keep) }) } func randomChatID(prefix string) (string, error) { value := make([]byte, 24) if _, err := rand.Read(value); err != nil { return "", err } return prefix + base64.RawURLEncoding.EncodeToString(value), nil } func validCorrelationID(value string) bool { if value == "" || len(value) > 512 { return false } for _, character := range value { if !(character >= 'a' && character <= 'z') && !(character >= 'A' && character <= 'Z') && !(character >= '0' && character <= '9') && !strings.ContainsRune("._:-", character) { return false } } return true } func validOpaqueProviderID(value string) bool { return strings.TrimSpace(value) != "" && len(value) <= 512 && !containsControl(value) } func containsControl(value string) bool { for _, character := range value { if unicode.IsControl(character) { return true } } return false }