456 lines
14 KiB
Go
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
|
|
}
|