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

456 lines
14 KiB
Go

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
}