This repository has been archived on 2026-04-03. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
core/server/handlers/completion.go
gobha b7746c3004 Changeset 0.37.14 (#226)
Co-authored-by: gobha <jasafpro@gmail.com>
Co-committed-by: gobha <jasafpro@gmail.com>
2026-03-23 16:47:48 +00:00

1972 lines
67 KiB
Go

package handlers
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"log"
"log/slog"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"chat-switchboard/auth"
capspkg "chat-switchboard/capabilities"
"chat-switchboard/crypto"
"chat-switchboard/database"
"chat-switchboard/events"
"chat-switchboard/filters"
"chat-switchboard/health"
"chat-switchboard/knowledge"
"chat-switchboard/metrics"
"chat-switchboard/models"
"chat-switchboard/notifications"
"chat-switchboard/providers"
"chat-switchboard/routing"
"chat-switchboard/sandbox"
"chat-switchboard/storage"
"chat-switchboard/store"
"chat-switchboard/tools"
)
// ── Request Types ───────────────────────────
type completionRequest struct {
ChannelID string `json:"channel_id"`
Content string `json:"content" binding:"required"`
Model string `json:"model,omitempty"`
PersonaID string `json:"persona_id,omitempty"`
ProviderConfigID string `json:"provider_config_id,omitempty"`
MaxTokens int `json:"max_tokens,omitempty"`
Temperature *float64 `json:"temperature,omitempty"`
TopP *float64 `json:"top_p,omitempty"`
Stream *bool `json:"stream,omitempty"`
FileIDs []string `json:"file_ids,omitempty"` // staged file UUIDs to include in request
DisabledTools []string `json:"disabled_tools,omitempty"` // tool names to exclude from this request
}
// CompletionHandler proxies LLM requests through the backend.
type CompletionHandler struct {
vault *crypto.KeyResolver
stores store.Stores
hub *events.Hub // WebSocket hub for browser tool bridge
objStore storage.ObjectStore // file storage for uploaded/generated content (nil = disabled)
embedder *knowledge.Embedder // for memory semantic recall (v0.18.0)
health HealthRecorder // provider health tracking (v0.22.0, nil = disabled)
healthStore HealthStatusQuerier // health status queries for routing (v0.22.2, nil = disabled)
router *routing.Evaluator // routing policy evaluator (v0.22.2, nil = disabled)
filterChain *filters.Chain // pre-completion filter chain (v0.29.0, nil = disabled)
runner *sandbox.Runner // v0.29.2: Starlark extension tool dispatch (nil = disabled)
}
// HealthRecorder is the interface for recording provider call outcomes.
// Matches health.Accumulator methods without importing the package.
type HealthRecorder interface {
RecordSuccess(providerConfigID string, latencyMs int)
RecordError(providerConfigID string, latencyMs int, errMsg string)
RecordTimeout(providerConfigID string, latencyMs int, errMsg string)
RecordRateLimit(providerConfigID string, latencyMs int, errMsg string) // v0.22.4
RecordToolSuccess(toolName string, latencyMs int) // v0.22.4
RecordToolError(toolName string, latencyMs int, errMsg string) // v0.22.4
}
// HealthStatusQuerier retrieves current provider health for routing decisions.
// Matches health.Store.ListAllCurrent without importing the package.
type HealthStatusQuerier interface {
ListAllCurrent(ctx context.Context) ([]models.ProviderHealthWindow, error)
}
// NewCompletionHandler creates a new handler.
func NewCompletionHandler(vault *crypto.KeyResolver, stores store.Stores, hub *events.Hub, objStore storage.ObjectStore, embedder *knowledge.Embedder) *CompletionHandler {
return &CompletionHandler{vault: vault, stores: stores, hub: hub, objStore: objStore, embedder: embedder}
}
// SetHealthRecorder attaches the provider health accumulator.
func (h *CompletionHandler) SetHealthRecorder(hr HealthRecorder) {
h.health = hr
}
// SetHealthStore attaches the health store for routing status queries (v0.22.2).
func (h *CompletionHandler) SetHealthStore(hs HealthStatusQuerier) {
h.healthStore = hs
}
// SetRoutingEvaluator attaches the routing policy engine (v0.22.2).
func (h *CompletionHandler) SetRoutingEvaluator(r *routing.Evaluator) {
h.router = r
}
// SetFilterChain attaches the pre-completion filter chain (v0.29.0).
func (h *CompletionHandler) SetFilterChain(fc *filters.Chain) {
h.filterChain = fc
}
// SetRunner attaches the Starlark sandbox runner for extension tool dispatch (v0.29.2).
func (h *CompletionHandler) SetRunner(r *sandbox.Runner) {
h.runner = r
}
// buildExtToolMap returns the extension tool map for the current user.
func (h *CompletionHandler) buildExtToolMap(ctx context.Context, userID string) map[string]*store.PackageRegistration {
return BuildExtToolMap(ctx, h.stores, userID)
}
// evaluateRouting applies routing policies to select the best provider config
// for this request. Returns the winning config details and a routing decision
// for observability. If routing is disabled or no policies match, returns
// the original resolved config unchanged.
func (h *CompletionHandler) evaluateRouting(
ctx context.Context,
userID, model, configID, providerID string,
) (winConfigID, winProviderID string, decision *routing.Decision) {
// Fallthrough: return original config if routing is disabled
if h.router == nil || h.stores.RoutingPolicies == nil {
return configID, providerID, nil
}
// Load active policies (scoped to user's teams)
var dbPolicies []models.RoutingPolicy
var err error
teamIDs, _ := h.stores.Teams.GetUserTeamIDs(ctx, userID)
if len(teamIDs) > 0 {
// Load global + team-scoped policies for user's teams
dbPolicies, err = h.stores.RoutingPolicies.ListActive(ctx)
} else {
dbPolicies, err = h.stores.RoutingPolicies.ListActive(ctx)
}
if err != nil || len(dbPolicies) == 0 {
return configID, providerID, nil
}
policies := routing.FromModels(dbPolicies)
// Build candidates from all configs accessible to this user
configs, err := h.stores.Providers.ListAccessible(ctx, userID)
if err != nil || len(configs) < 2 {
// No point routing with a single config
return configID, providerID, nil
}
candidates := make([]routing.Candidate, 0, len(configs))
for _, cfg := range configs {
if !cfg.IsActive {
continue
}
candidates = append(candidates, routing.Candidate{
ConfigID: cfg.ID,
ProviderID: cfg.Provider,
Model: model,
Endpoint: cfg.Endpoint,
})
}
if len(candidates) < 2 {
return configID, providerID, nil
}
// Build health status map
healthStatus := make(map[string]models.ProviderStatus)
if h.healthStore != nil {
windows, _ := h.healthStore.ListAllCurrent(ctx)
for _, w := range windows {
healthStatus[w.ProviderConfigID] = health.DeriveStatus(w.ErrorRate(), w.RequestCount)
}
}
routingCtx := &routing.Context{
UserID: userID,
TeamIDs: teamIDs,
Model: model,
ConfigID: configID,
HealthStatus: healthStatus,
Pricing: map[string]*models.ModelPricing{}, // populated when cost_limit routing is enabled (deferred)
}
ranked, dec := h.router.Evaluate(routingCtx, policies, candidates)
if len(ranked) > 0 && ranked[0].ConfigID != configID {
return ranked[0].ConfigID, ranked[0].ProviderID, dec
}
return configID, providerID, dec
}
// ── Chat Completion ─────────────────────────
// POST /api/v1/chat/completions
//
// Flow:
// 1. Validate request, verify chat ownership
// 2. Resolve api_config (from request, chat, or user default)
// 3. Load conversation history from messages
// 4. Persist user message
// 5. Attach tool definitions if model supports tools
// 6. Call provider (stream or non-stream)
// 7. Tool loop: if LLM requests tool calls, execute them and re-call
// 8. Stream SSE to client / return JSON
// 9. Persist assistant message with token counts
func (h *CompletionHandler) Complete(c *gin.Context) {
var req completionRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
channelID := req.ChannelID
if channelID == "" {
// v0.24.3: Workflow API routes have channel in URL param
channelID = c.Param("id")
}
if channelID == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "channel_id is required"})
return
}
userID := getUserID(c)
// v0.33.0: Correlation ID — same as request_id, propagated through
// completion chain for cross-referencing logs.
correlationID, _ := c.Get("request_id")
slog.Info("completion.start",
"request_id", correlationID,
"channel_id", channelID,
"user_id", userID,
"model", req.Model,
)
// v0.24.3: Session participants are pre-validated by AuthOrSession middleware
if isSessionAuth(c) {
if !sessionCanAccessChannel(c, channelID) {
c.JSON(http.StatusForbidden, gin.H{"error": "session not authorized for this channel"})
return
}
// Use session_id as effective identity for cursor/participant tracking.
// persistMessage will record participant_type=user with this ID —
// acceptable for v0.24.3; v0.25.0 can refine to participant_type=session.
userID = c.GetString("session_id")
} else {
// Verify participation: check channel_participants first, fall back to
// legacy channels.user_id ownership for direct channels.
isParticipant, _ := h.stores.Channels.IsParticipant(c.Request.Context(), channelID, "user", userID)
if !isParticipant {
if !userOwnsChannel(c, channelID, userID) {
return
}
}
}
// ── ai_mode guard (v0.23.1) ────────────────────────────────────────────────
// Check channel ai_mode before doing any work. For DM channels with
// mention_only mode and no @mention in the content, persist the message
// but skip the completion entirely.
{
var aiMode string
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COALESCE(ai_mode, 'auto') FROM channels WHERE id = $1
`), channelID).Scan(&aiMode)
if aiMode == "off" || (aiMode == "mention_only" && extractFirstMention(req.Content) == "") {
// Persist the user message before returning
msgID, err := h.persistMessage(channelID, userID, "user", req.Content, "", 0, 0, nil, nil, "", "")
if err != nil {
log.Printf("[ai_skip] Failed to persist user message: %v", err)
}
if h.hub != nil && msgID != "" {
broadcastUserMessage(c.Request.Context(), h.hub, channelID, msgID, req.Content, userID)
}
c.JSON(http.StatusOK, gin.H{"status": "delivered", "ai_skipped": true, "message_id": msgID})
return
}
}
// ── Persona unwrap: persona overrides defaults, explicit request fields win ──
var personaSystemPrompt string
var personaID string // tracks active persona for KB scoping
var personaThinkingBudget *int // persona-level thinking budget for hook injection (v0.22.1)
// v0.25.0: Query channel metadata once — used by group leader check, tool context,
// and execution context. Avoids redundant SELECTs on channels.
var channelType string
var channelTeamID *string
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COALESCE(type, 'direct'), team_id FROM channels WHERE id = $1
`), channelID).Scan(&channelType, &channelTeamID)
teamID := ""
if channelTeamID != nil {
teamID = *channelTeamID
}
// v0.23.2: Group leader default — when type=group and no @mention, use the leader persona
if req.PersonaID == "" && extractFirstMention(req.Content) == "" {
if channelType == "group" {
// Find the leader persona via channel_participants → persona_group_members
var leaderPersonaID string
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT cp.participant_id
FROM channel_participants cp
JOIN persona_group_members pgm ON pgm.persona_id = cp.participant_id
WHERE cp.channel_id = $1
AND cp.participant_type = 'persona'
AND pgm.is_leader = true
LIMIT 1
`), channelID).Scan(&leaderPersonaID)
// Fallback: first persona participant if no leader flag set
if leaderPersonaID == "" {
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'persona'
ORDER BY created_at LIMIT 1
`), channelID).Scan(&leaderPersonaID)
}
if leaderPersonaID != "" {
req.PersonaID = leaderPersonaID
log.Printf("[group] channel %s: leader persona %s", channelID[:min(8, len(channelID))], leaderPersonaID[:min(8, len(leaderPersonaID))])
}
}
}
if req.PersonaID != "" {
persona := ResolvePersona(h.stores, req.PersonaID, userID)
if persona == nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "persona not found or not accessible"})
return
}
personaID = persona.ID
// Persona provides defaults; explicit request fields take priority
if req.Model == "" {
req.Model = persona.BaseModelID
}
if req.ProviderConfigID == "" && persona.ProviderConfigID != nil {
req.ProviderConfigID = *persona.ProviderConfigID
}
if req.Temperature == nil && persona.Temperature != nil {
req.Temperature = persona.Temperature
}
if req.MaxTokens == 0 && persona.MaxTokens != nil {
req.MaxTokens = *persona.MaxTokens
}
if persona.SystemPrompt != "" {
personaSystemPrompt = persona.SystemPrompt
}
personaThinkingBudget = persona.ThinkingBudget
}
// ── Project persona fallback (v0.19.2): if no explicit persona, check project ──
if req.PersonaID == "" && h.stores.Projects != nil {
projID, _ := h.stores.Projects.GetProjectIDForChannel(context.Background(), channelID)
if projID != "" {
if proj, err := h.stores.Projects.GetByID(context.Background(), projID); err == nil {
if pid, ok := proj.Settings["persona_id"].(string); ok && pid != "" {
if persona := ResolvePersona(h.stores, pid, userID); persona != nil {
personaID = persona.ID
if req.Model == "" {
req.Model = persona.BaseModelID
}
if req.ProviderConfigID == "" && persona.ProviderConfigID != nil {
req.ProviderConfigID = *persona.ProviderConfigID
}
if req.Temperature == nil && persona.Temperature != nil {
req.Temperature = persona.Temperature
}
if req.MaxTokens == 0 && persona.MaxTokens != nil {
req.MaxTokens = *persona.MaxTokens
}
if persona.SystemPrompt != "" {
personaSystemPrompt = persona.SystemPrompt
}
personaThinkingBudget = persona.ThinkingBudget
}
}
}
}
}
// Resolve provider config
providerCfg, providerID, model, configID, providerScope, err := h.resolveConfig(userID, channelID, req)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
// ── Routing policy evaluation (v0.22.2) ──────────
// After resolving the primary config, evaluate routing policies to
// potentially select a different (better/healthier) provider.
var routingDecision *routing.Decision
if h.router != nil {
winConfigID, _, dec := h.evaluateRouting(c.Request.Context(), userID, model, configID, providerID)
routingDecision = dec
if winConfigID != configID {
// Routing selected a different provider — reload its credentials
req.ProviderConfigID = winConfigID
providerCfg2, providerID2, model2, configID2, providerScope2, err := h.resolveConfig(userID, channelID, req)
if err == nil {
providerCfg = providerCfg2
providerID = providerID2
if model2 != "" {
model = model2
}
configID = configID2
providerScope = providerScope2
} else {
// Routing target failed to load — fall back to original.
log.Printf("[routing] Failed to load winning config %s: %v — using original %s", winConfigID, err, configID)
}
}
}
// ── Inject persona-level thinking budget into provider settings (v0.22.1) ──
// When a persona specifies a thinking budget, promote it into provider
// settings so hooks can activate extended thinking automatically.
if personaThinkingBudget != nil && *personaThinkingBudget > 0 {
if providerCfg.Settings == nil {
providerCfg.Settings = make(map[string]interface{})
}
providerCfg.Settings["extended_thinking"] = true
providerCfg.Settings["thinking_budget"] = *personaThinkingBudget
}
// ── Team policy: require_private_providers ──
if err := enforcePrivateProviderPolicy(c.Request.Context(), h.stores, userID, configID); err != nil {
c.JSON(http.StatusForbidden, gin.H{"error": err.Error()})
return
}
provider, err := providers.Get(providerID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
// Routing observability header (v0.22.2)
c.Header("X-Switchboard-Provider", providerID+"/"+configID)
if routingDecision != nil {
c.Header("X-Switchboard-Route", routingDecision.PolicyName)
if routingDecision.FallbackDepth > 0 {
c.Header("X-Switchboard-Fallback", fmt.Sprintf("%d", routingDecision.FallbackDepth))
}
}
// Load conversation history
messages, err := h.loadConversation(channelID, userID, personaSystemPrompt, personaID, model)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load conversation"})
return
}
// ── Workflow form template injection (v0.26.4) ──
// When the channel is a workflow instance, inject the current stage's
// form_template as a system message so the persona knows what to collect.
if channelType == "workflow" {
if hint := h.buildWorkflowFormHint(c.Request.Context(), channelID); hint != "" {
messages = append(messages[:1], append([]providers.Message{{
Role: "system",
Content: hint,
}}, messages[1:]...)...)
}
}
// Resolve capabilities early — needed for vision gating below
caps := h.getModelCapabilities(c, model, configID)
// Build user message — multimodal if files are present
userMsg := providers.Message{
Role: "user",
Content: req.Content,
}
if len(req.FileIDs) > 0 && h.objStore != nil {
parts, augContent, _, err := h.buildMultimodalParts(c, channelID, req.Content, req.FileIDs, caps)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if len(parts) > 0 {
// Multimodal (images present): use ContentParts
userMsg.ContentParts = parts
} else if augContent != req.Content {
// Doc-only: use augmented text content
userMsg.Content = augContent
}
}
messages = append(messages, userMsg)
// Persist user message (text-only for storage — multimodal parts are ephemeral)
msgID, err := h.persistMessage(channelID, userID, "user", req.Content, "", 0, 0, nil, nil, "", "")
if err != nil {
log.Printf("Failed to persist user message: %v", err)
}
// Broadcast user message to all channel participants
if h.hub != nil && msgID != "" {
broadcastUserMessage(c.Request.Context(), h.hub, channelID, msgID, req.Content, userID)
}
// Link files to the persisted message
if msgID != "" && len(req.FileIDs) > 0 {
for _, fID := range req.FileIDs {
if err := h.stores.Files.SetMessageID(c.Request.Context(), fID, msgID); err != nil {
log.Printf("Failed to link file %s to message %s: %v", fID, msgID, err)
}
}
}
// ── Workspace resolution (v0.21.1) ────────
// Resolve effective workspace: channel.workspace_id > project.workspace_id
workspaceID, _ := h.stores.Channels.ResolveWorkspaceID(c.Request.Context(), channelID)
// ── Tool context (v0.25.0) ─────────────────
// Uses channel metadata queried above (channelType, teamID).
tctx := tools.ToolContext{
ChannelType: channelType,
WorkspaceID: workspaceID,
TeamID: teamID,
IsVisitor: isSessionAuth(c),
// PersonaID set below after persona resolution
}
// ── @mention routing (v0.23.0 / v0.23.1) ─────────────────────────────────
// Resolve @persona-handle, @model-id, or @username.
// - User mention → skip completion, notify recipient (v0.23.1)
// - AI mention → override model/persona for this turn (v0.23.0)
mentionModel, mentionConfig, mentionPersona, mentionedUserID := h.resolveMention(c.Request.Context(), userID, req.Content)
// v0.23.2: @all fan-out — route to every persona participant (depth-1 only)
if strings.Contains(strings.ToLower(req.Content), "@all") {
// Get all persona participants for this channel
allRows, err := database.DB.QueryContext(c.Request.Context(), database.Q(`
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'persona'
ORDER BY created_at
`), channelID)
if err == nil {
defer allRows.Close()
var allPersonaIDs []string
for allRows.Next() {
var pid string
if allRows.Scan(&pid) == nil {
allPersonaIDs = append(allPersonaIDs, pid)
}
}
if len(allPersonaIDs) > 0 {
// Use the first persona for the streamed response
first := ResolvePersona(h.stores, allPersonaIDs[0], userID)
if first != nil {
mentionModel = first.BaseModelID
if first.ProviderConfigID != nil {
mentionConfig = *first.ProviderConfigID
}
mentionPersona = first
}
// Chain remaining personas via goroutine after the stream completes
if len(allPersonaIDs) > 1 && h.hub != nil {
remaining := allPersonaIDs[1:]
go func() {
for _, pid := range remaining {
time.Sleep(500 * time.Millisecond)
h.chainToPersona(channelID, userID, pid, req.Content, 1)
}
}()
}
}
}
}
if mentionedUserID != "" {
// Human @mention — message already persisted; notify recipient, skip AI.
if h.hub != nil {
payload, _ := json.Marshal(map[string]any{
"channel_id": channelID,
"from_user": userID,
"content": truncateContent(req.Content, 120),
})
h.hub.PublishToUser(mentionedUserID, events.Event{
Label: "user.mentioned",
Payload: payload,
Ts: time.Now().UnixMilli(),
})
}
// Persist notification for bell/inbox (v0.28.2)
if svc := notifications.Default(); svc != nil {
notifications.NotifyUserMentioned(svc, mentionedUserID, channelID, userID, truncateContent(req.Content, 120))
}
c.JSON(http.StatusOK, gin.H{"status": "delivered", "mentioned_user": mentionedUserID})
return
}
if mentionModel != "" {
req.Model = mentionModel
if mentionConfig != "" {
req.ProviderConfigID = mentionConfig
}
if mentionPersona != nil {
personaID = mentionPersona.ID
personaThinkingBudget = mentionPersona.ThinkingBudget
if mentionPersona.SystemPrompt != "" {
if len(messages) > 0 && messages[0].Role == "system" {
messages[0].Content = mentionPersona.SystemPrompt
} else {
messages = append([]providers.Message{{Role: "system", Content: mentionPersona.SystemPrompt}}, messages...)
}
}
// Context boundary: tell this persona to ignore other personas' styles
boundary := fmt.Sprintf(
"[You are now %s. Previous assistant messages may be from different AI personas with different styles — ignore their tone, character, and mannerisms. Respond only as yourself per your system prompt.]",
mentionPersona.Name,
)
if len(messages) >= 2 {
last := messages[len(messages)-1]
messages[len(messages)-1] = providers.Message{Role: "system", Content: boundary}
messages = append(messages, last)
}
}
// Re-resolve provider config for the @mentioned target
providerCfg, providerID, model, configID, providerScope, err = h.resolveConfig(userID, channelID, req)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
caps = h.getModelCapabilities(c, model, configID)
if personaThinkingBudget != nil && *personaThinkingBudget > 0 {
if providerCfg.Settings == nil {
providerCfg.Settings = make(map[string]interface{})
}
providerCfg.Settings["extended_thinking"] = true
providerCfg.Settings["thinking_budget"] = *personaThinkingBudget
}
provider, err = providers.Get(providerID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
} else {
// No @mention: inject a boundary if conversation has persona responses
// to prevent the default model from mimicking persona styles.
if h.conversationHasPersonaMessages(c.Request.Context(), channelID) {
boundary := "[Previous assistant messages in this conversation may include responses from AI personas with specific characters and styles. You are the default assistant — respond in your own natural style, not theirs.]"
if len(messages) >= 2 {
last := messages[len(messages)-1]
messages[len(messages)-1] = providers.Message{Role: "system", Content: boundary}
messages = append(messages, last)
}
}
}
// Build provider request
provReq := providers.CompletionRequest{
Model: model,
Messages: messages,
}
if req.MaxTokens > 0 {
provReq.MaxTokens = req.MaxTokens
} else {
// ResolveMaxOutput checks: caps → known models → context/8 → 4096
provReq.MaxTokens = capspkg.ResolveMaxOutput(model, caps)
}
if req.Temperature != nil {
provReq.Temperature = req.Temperature
}
if req.TopP != nil {
provReq.TopP = req.TopP
}
// Attach tool definitions if model supports tool calling and tools are available
tctx.PersonaID = personaID // finalize after all persona resolution (@mention, group leader, etc.)
hasBrowserTools := h.hub != nil && h.hub.IsConnected(userID)
if caps.ToolCalling && (tools.HasTools() || hasBrowserTools) {
provReq.Tools = h.buildToolDefs(c.Request.Context(), userID, hasBrowserTools, req.DisabledTools, tctx, personaID)
}
// ── Permission pre-flight ─────────────────────────────────────────────
// Budget enforcement and model allowlist only apply when the request draws
// from a team or global provider. BYOK (personal scope) is the user's money
// — no ceiling applies.
// v0.24.3: Session participants bypass user-level budget/allowlist checks.
// They use the workflow channel's allocated resources.
role, _ := c.Get("role")
if role != "admin" && providerScope != "personal" && !isSessionAuth(c) {
// Token budget
if h.stores.Usage != nil {
budget, err := auth.ResolveTokenBudget(c.Request.Context(), h.stores, userID)
if err == nil && budget != nil {
var opts store.UsageQueryOptions
now := time.Now()
if budget.Period == "daily" {
dayStart := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC)
opts.Since = &dayStart
} else {
monthStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC)
opts.Since = &monthStart
}
totals, err := h.stores.Usage.GetPersonalTotals(c.Request.Context(), userID, opts)
if err == nil && totals != nil {
used := int64(totals.InputTokens + totals.OutputTokens)
if used >= budget.Limit {
c.JSON(http.StatusTooManyRequests, gin.H{"error": "token budget exceeded"})
return
}
}
}
}
// Model allowlist
allowlist, err := auth.ResolveModelAllowlist(c.Request.Context(), h.stores, userID)
if err == nil && allowlist != nil && !allowlist[model] {
c.JSON(http.StatusForbidden, gin.H{"error": "model not permitted for your account: " + model})
return
}
}
// Determine streaming
stream := true // default
if req.Stream != nil {
stream = *req.Stream
}
if stream {
h.streamCompletion(c, provider, providerCfg, provReq, channelID, userID, model, providerID, configID, providerScope, personaID, workspaceID, teamID)
} else {
h.syncCompletion(c, provider, providerCfg, provReq, channelID, userID, model, providerID, configID, providerScope, personaID, workspaceID, teamID)
}
}
// ── Multi-Model Sequential Stream (v0.20.0) ──
// multiModelStream handles the case where a message @mentions multiple
// channel models. It streams each model's response sequentially within
// a single SSE response, sending model_start/model_end delimiters.
func (h *CompletionHandler) multiModelStream(
c *gin.Context,
targets []models.ChannelModel,
messages []providers.Message,
channelID, userID, personaID, personaSystemPrompt, workspaceID, teamID string,
req completionRequest,
) {
// Set SSE headers once for the entire multi-model stream
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no")
c.Status(http.StatusOK)
flusher, _ := c.Writer.(http.Flusher)
flush := func() {
if flusher != nil {
flusher.Flush()
}
}
sendSSE := func(data string) {
fmt.Fprintf(c.Writer, "data: %s\n\n", data)
flush()
}
sendEvent := func(event, data string) {
fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", event, data)
flush()
}
for _, target := range targets {
// Notify client which model is about to respond
startJSON, _ := json.Marshal(map[string]string{
"model_id": target.ModelID,
"display_name": target.DisplayName,
})
sendEvent("model_start", string(startJSON))
// Resolve config for this specific model
targetReq := req
targetReq.Model = target.ModelID
if target.ProviderConfigID != "" {
targetReq.ProviderConfigID = target.ProviderConfigID
}
providerCfg, providerID, model, configID, providerScope, err := h.resolveConfig(userID, channelID, targetReq)
if err != nil {
sendSSE(fmt.Sprintf(`{"error":"failed to resolve config for %s: %s"}`, target.DisplayName, escapeJSON(err.Error())))
sendEvent("model_end", string(startJSON))
continue
}
provider, err := providers.Get(providerID)
if err != nil {
sendSSE(fmt.Sprintf(`{"error":"provider unavailable for %s"}`, target.DisplayName))
sendEvent("model_end", string(startJSON))
continue
}
// Build per-model messages (inject model-specific system prompt if set)
modelMessages := make([]providers.Message, len(messages))
copy(modelMessages, messages)
if target.SystemPrompt != "" {
// Prepend model-specific system prompt
modelMessages = append([]providers.Message{{
Role: "system",
Content: target.SystemPrompt,
}}, modelMessages...)
}
caps := h.getModelCapabilities(c, model, configID)
provReq := providers.CompletionRequest{
Model: model,
Messages: modelMessages,
}
if targetReq.MaxTokens > 0 {
provReq.MaxTokens = targetReq.MaxTokens
} else {
provReq.MaxTokens = capspkg.ResolveMaxOutput(model, caps)
}
if targetReq.Temperature != nil {
provReq.Temperature = targetReq.Temperature
}
if targetReq.TopP != nil {
provReq.TopP = targetReq.TopP
}
hasBrowserTools := h.hub != nil && h.hub.IsConnected(userID)
if caps.ToolCalling && (tools.HasTools() || hasBrowserTools) {
tctx := tools.ToolContext{
WorkspaceID: workspaceID,
PersonaID: personaID,
IsVisitor: isSessionAuth(c),
}
provReq.Tools = h.buildToolDefs(c.Request.Context(), userID, hasBrowserTools, req.DisabledTools, tctx, personaID)
}
// Apply provider-specific request hooks (v0.22.1)
if hooks := providers.GetHooks(providerID); hooks != nil {
hooks.PreRequest(providerCfg, &provReq)
}
// Stream this model's response using the shared streaming core
multiExtTools := h.buildExtToolMap(c.Request.Context(), userID)
result := streamModelResponse(c, provider, providerCfg, &provReq, model, providerID, target.DisplayName, userID, channelID, personaID, workspaceID, configID, teamID, h.hub, h.health, h.runner, multiExtTools)
// Persist assistant message with model attribution
if result.Content != "" {
pID := ""
if target.PersonaID != nil {
pID = *target.PersonaID
}
if _, err := h.persistMessage(channelID, userID, "assistant", result.Content, model, result.InputTokens, result.OutputTokens, nil, toolActivityJSON(result.ToolActivity), configID, pID); err != nil {
log.Printf("Failed to persist multi-model assistant message: %v", err)
}
}
// Log usage
h.logUsage(c, channelID, userID, configID, providerScope, model,
result.InputTokens, result.OutputTokens,
result.CacheCreationTokens, result.CacheReadTokens)
sendEvent("model_end", string(startJSON))
}
sendSSE("[DONE]")
}
// buildToolDefs delegates to the standalone BuildToolDefs function.
// See resolve.go for the full implementation.
func (h *CompletionHandler) buildToolDefs(ctx context.Context, userID string, includeBrowser bool, disabledTools []string, tctx tools.ToolContext, personaID string) []providers.ToolDef {
return BuildToolDefs(ctx, h.stores, userID, includeBrowser, disabledTools, tctx, personaID)
}
// ── List Available Tools ────────────────────
// GET /api/v1/tools
//
// Returns all registered server-side tools plus browser extension tools
// for the current user. Used by the frontend to build the tools toggle menu.
type toolInfo struct {
Name string `json:"name"`
DisplayName string `json:"display_name"`
Description string `json:"description"`
Category string `json:"category"`
}
func (h *CompletionHandler) ListTools(c *gin.Context) {
userID := c.GetString("user_id")
// Server-side tools
allDefs := tools.AllDefinitions()
result := make([]toolInfo, 0, len(allDefs))
for _, d := range allDefs {
result = append(result, toolInfo{
Name: d.Name,
DisplayName: d.DisplayName,
Description: d.Description,
Category: d.Category,
})
}
// Browser extension tools
if h.hub != nil && h.hub.IsConnected(userID) {
pkgs, err := h.stores.Packages.ListForUser(c.Request.Context(), userID)
if err == nil {
for _, pkg := range pkgs {
if pkg.Tier != "browser" {
continue
}
var manifest struct {
Tools []struct {
Name string `json:"name"`
DisplayName string `json:"display_name"`
Description string `json:"description"`
} `json:"tools"`
}
if err := json.Unmarshal(marshalManifest(pkg.Manifest), &manifest); err != nil {
continue
}
for _, t := range manifest.Tools {
displayName := t.DisplayName
if displayName == "" {
displayName = t.Name
}
result = append(result, toolInfo{
Name: t.Name,
DisplayName: displayName,
Description: t.Description,
Category: "browser",
})
}
}
}
}
c.JSON(http.StatusOK, gin.H{"data": result})
}
// ── Streaming Completion (SSE) with Tool Loop ──
func (h *CompletionHandler) streamCompletion(
c *gin.Context,
provider providers.Provider,
cfg providers.ProviderConfig,
req providers.CompletionRequest,
channelID, userID, model, providerID, configID, providerScope, personaID, workspaceID, teamID string,
) {
// Apply provider-specific request hooks (v0.22.1)
if hooks := providers.GetHooks(providerID); hooks != nil {
hooks.PreRequest(cfg, &req)
}
extTools := h.buildExtToolMap(c.Request.Context(), userID)
result := streamWithToolLoop(c, provider, cfg, &req, model, providerID, userID, channelID, personaID, workspaceID, configID, teamID, h.hub, h.health, h.runner, extTools)
// Persist assistant response
if result.Content != "" {
asstMsgID, err := h.persistMessage(channelID, userID, "assistant", result.Content, model, result.InputTokens, result.OutputTokens, nil, toolActivityJSON(result.ToolActivity), configID, personaID)
if err != nil {
log.Printf("Failed to persist assistant message: %v", err)
}
// Broadcast assistant message to other channel participants
if h.hub != nil && asstMsgID != "" {
broadcastAssistantMessage(c.Request.Context(), h.hub, channelID, asstMsgID, result.Content, model, userID, personaID)
}
// AI-to-AI chaining: if the response @mentions another persona participant,
// trigger a follow-up completion asynchronously via WebSocket delivery.
go func() {
defer func() {
if r := recover(); r != nil {
log.Printf("[chain] PANIC recovered: %v", r)
}
}()
h.chainIfMentioned(channelID, userID, personaID, result.Content, 0)
}()
}
// Log usage
h.logUsage(c, channelID, userID, configID, providerScope, model,
result.InputTokens, result.OutputTokens,
result.CacheCreationTokens, result.CacheReadTokens)
}
// ── Non-Streaming Completion with Tool Loop ──
func (h *CompletionHandler) syncCompletion(
c *gin.Context,
provider providers.Provider,
cfg providers.ProviderConfig,
req providers.CompletionRequest,
channelID, userID, model, providerID, configID, providerScope, personaID, workspaceID, teamID string,
) {
// Apply provider-specific request hooks (v0.22.1)
if hooks := providers.GetHooks(providerID); hooks != nil {
hooks.PreRequest(cfg, &req)
}
extTools := h.buildExtToolMap(c.Request.Context(), userID)
result := CoreToolLoop(c.Request.Context(), LoopConfig{
Provider: provider,
Cfg: cfg,
Req: &req,
Model: model,
ProviderType: providerID,
ExecCtx: tools.ExecutionContext{
UserID: userID,
ChannelID: channelID,
PersonaID: personaID,
WorkspaceID: workspaceID,
TeamID: teamID,
},
Hub: h.hub,
Health: h.health,
ConfigID: configID,
Budget: LoopBudget{},
Streaming: false,
Runner: h.runner,
ExtTools: extTools,
}, accumSink{})
// Provider error
if result.Error != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": "provider error: " + result.Error.Error()})
return
}
// Max iterations hit with no final content
if result.BudgetExceeded == "max_rounds" && result.Content == "" {
c.JSON(http.StatusOK, gin.H{
"model": model,
"choices": []gin.H{
{
"message": gin.H{
"role": "assistant",
"content": "[Tool execution limit reached]",
},
"finish_reason": "stop",
},
},
})
return
}
finalContent := result.Content
// Persist assistant response
if _, err := h.persistMessage(channelID, userID, "assistant", finalContent, model, result.InputTokens, result.OutputTokens, nil, toolActivityJSON(result.ToolActivity), configID, personaID); err != nil {
log.Printf("Failed to persist assistant message: %v", err)
}
// AI-to-AI chaining
go func() {
defer func() {
if r := recover(); r != nil {
log.Printf("[chain] PANIC recovered: %v", r)
}
}()
h.chainIfMentioned(channelID, userID, personaID, finalContent, 0)
}()
// Log usage
h.logUsage(c, channelID, userID, configID, providerScope, model,
result.InputTokens, result.OutputTokens,
result.CacheCreationTokens, result.CacheReadTokens)
// Return OpenAI-compatible response
finishReason := "stop"
if result.BudgetExceeded != "" {
finishReason = "budget_exceeded"
}
c.JSON(http.StatusOK, gin.H{
"model": model,
"choices": []gin.H{
{
"message": gin.H{
"role": "assistant",
"content": finalContent,
},
"finish_reason": finishReason,
},
},
"usage": gin.H{
"prompt_tokens": result.InputTokens,
"completion_tokens": result.OutputTokens,
"total_tokens": result.InputTokens + result.OutputTokens,
},
})
}
// ── Model Capabilities ──────────────────────
// escapeJSON escapes a string for safe embedding in a JSON string value.
func escapeJSON(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `"`, `\"`)
s = strings.ReplaceAll(s, "\n", `\n`)
s = strings.ReplaceAll(s, "\r", `\r`)
s = strings.ReplaceAll(s, "\t", `\t`)
return s
}
// getModelCapabilities looks up capabilities from model_catalog DB,
// then overlays with known model defaults and heuristic detection.
func (h *CompletionHandler) getModelCapabilities(c *gin.Context, model, apiConfigID string) models.ModelCapabilities {
return ResolveModelCaps(h.stores.Catalog, model, apiConfigID)
}
// ── Multimodal Assembly ─────────────────────
// Builds content parts from files for the user message.
//
// Returns:
// - parts: ContentParts array (non-nil only when images are present)
// - augContent: enriched text content with document context (for doc-only case)
// - validIDs: file IDs that were successfully processed
// - error: if any file is invalid or vision is needed but missing
//
// Rules:
// - Images → base64 data URI (requires vision capability)
// - Documents with extracted_text → text injection
// - Documents without extraction → filename placeholder
// - Text is always the first part
func (h *CompletionHandler) buildMultimodalParts(
c *gin.Context,
channelID, textContent string,
fileIDs []string,
caps models.ModelCapabilities,
) ([]providers.ContentPart, string, []string, error) {
parts := []providers.ContentPart{
{Type: "text", Text: textContent},
}
var docTexts []string
var validIDs []string
hasImage := false
for _, fileID := range fileIDs {
att, err := h.stores.Files.GetByID(c.Request.Context(), fileID)
if err != nil {
return nil, "", nil, fmt.Errorf("file %s not found", fileID)
}
// Security: verify file belongs to this channel
if att.ChannelID != channelID {
return nil, "", nil, fmt.Errorf("file %s does not belong to this channel", fileID)
}
if isImageContentType(att.ContentType) {
// Vision gating: reject images if model lacks vision
if !caps.Vision {
return nil, "", nil, fmt.Errorf("model does not support image input; remove image files or choose a vision-capable model")
}
// Read image from storage, base64 encode, build data URI
reader, _, _, err := h.objStore.Get(c.Request.Context(), att.StorageKey)
if err != nil {
log.Printf("Failed to read file %s from storage: %v", fileID, err)
parts = append(parts, providers.ContentPart{
Type: "text",
Text: fmt.Sprintf("[Image: %s — failed to read from storage]", att.Filename),
})
validIDs = append(validIDs, fileID)
continue
}
data, err := io.ReadAll(reader)
reader.Close()
if err != nil {
log.Printf("Failed to read file %s bytes: %v", fileID, err)
validIDs = append(validIDs, fileID)
continue
}
b64 := base64.StdEncoding.EncodeToString(data)
dataURI := fmt.Sprintf("data:%s;base64,%s", att.ContentType, b64)
parts = append(parts, providers.ContentPart{
Type: "image_url",
ImageURL: &providers.ImageURL{
URL: dataURI,
Detail: "auto",
},
})
hasImage = true
} else if att.ExtractedText != nil && *att.ExtractedText != "" {
// Document with extracted text → inject as context
docText := fmt.Sprintf("[Document: %s]\n%s", att.Filename, *att.ExtractedText)
parts = append(parts, providers.ContentPart{
Type: "text",
Text: docText,
})
docTexts = append(docTexts, docText)
} else {
// Document without extraction (pending, failed, or not extractable)
status := "pending"
if s, ok := att.Metadata["extraction_status"].(string); ok {
status = s
}
placeholder := fmt.Sprintf("[Attached file: %s (extraction %s)]", att.Filename, status)
parts = append(parts, providers.ContentPart{
Type: "text",
Text: placeholder,
})
docTexts = append(docTexts, placeholder)
}
validIDs = append(validIDs, fileID)
}
// If images present → use ContentParts (multimodal array)
if hasImage {
return parts, textContent, validIDs, nil
}
// Doc-only: merge document context into a single text string (more efficient)
augmented := textContent
for _, dt := range docTexts {
augmented += "\n\n" + dt
}
return nil, augmented, validIDs, nil
}
// isImageContentType returns true for MIME types that should be sent as
// base64 image content parts rather than text extraction.
func isImageContentType(ct string) bool {
return strings.HasPrefix(ct, "image/")
}
// ── @mention Resolution (v0.23.0) ───────────
// conversationHasPersonaMessages checks if any assistant message in the
// channel was generated by a persona (participant_type = 'persona').
// Used to decide whether to inject a context boundary for the default model.
func (h *CompletionHandler) conversationHasPersonaMessages(ctx context.Context, channelID string) bool {
var id string
err := database.DB.QueryRowContext(ctx, database.Q(`
SELECT id FROM messages
WHERE channel_id = $1 AND role = 'assistant' AND participant_type = 'persona'
LIMIT 1
`), channelID).Scan(&id)
return err == nil && id != ""
}
// buildWorkflowFormHint loads the current workflow stage's form_template
// and returns a system message instructing the persona what to collect.
// Returns empty string if not a workflow or no template defined.
// v0.29.3: migrated from raw database.DB to stores, stage_mode aware.
func (h *CompletionHandler) buildWorkflowFormHint(ctx context.Context, channelID string) string {
ws, err := h.stores.Channels.GetWorkflowStatus(ctx, channelID)
if err != nil || ws == nil || ws.WorkflowID == nil || *ws.WorkflowID == "" {
return ""
}
if h.stores.Workflows == nil {
return ""
}
stages, err := h.stores.Workflows.ListStages(ctx, *ws.WorkflowID)
if err != nil || ws.CurrentStage >= len(stages) {
return ""
}
stage := stages[ws.CurrentStage]
// v0.29.3: form_only stages don't use LLM — no hint needed
if stage.StageMode == models.StageModeFormOnly {
return ""
}
if len(stage.FormTemplate) == 0 || string(stage.FormTemplate) == "{}" {
return ""
}
// v0.29.3: Try typed form template first
if tpl := models.ParseTypedFormTemplate(stage.FormTemplate); tpl != nil {
hint := "You are in workflow stage: " + stage.Name + ".\n\n"
if stage.StageMode == models.StageModeFormChat {
hint += "A form has been provided for the visitor to fill in structured data. " +
"You can assist them with questions about the form fields or discuss related topics.\n\n" +
"Form fields:\n"
} else {
hint += "Collect the following information from the user through natural conversation:\n\n"
}
for _, f := range tpl.Fields {
hint += "- " + f.Label
if f.Required {
hint += " (required)"
}
hint += "\n"
}
if stage.StageMode != models.StageModeFormChat {
hint += "\nWhen you have gathered all required information, call the workflow_advance tool " +
"with the collected data as a JSON object. Do not advance until all required fields are filled."
}
return hint
}
// Legacy format: map[string]interface{} with {description, required}
var fields map[string]interface{}
if err := json.Unmarshal(stage.FormTemplate, &fields); err != nil {
return ""
}
if len(fields) == 0 {
return ""
}
hint := "You are in workflow stage: " + stage.Name + ".\n\n"
hint += "Collect the following information from the user through natural conversation:\n\n"
for name, spec := range fields {
switch v := spec.(type) {
case map[string]interface{}:
desc, _ := v["description"].(string)
required, _ := v["required"].(bool)
if desc != "" {
hint += "- " + name + ": " + desc
} else {
hint += "- " + name
}
if required {
hint += " (required)"
}
hint += "\n"
default:
hint += "- " + name + "\n"
}
}
hint += "\nWhen you have gathered all required information, call the workflow_advance tool " +
"with the collected data as a JSON object. Do not advance until all required fields are filled."
return hint
}
// buildParticipantHint builds a system message listing available @mentionable
// personas so the LLM knows who it can direct responses to.
// Only includes personas accessible to this user. Excludes the current persona.
func (h *CompletionHandler) buildParticipantHint(channelID, userID, currentPersonaID string) string {
ctx := context.Background()
personas, err := h.stores.Personas.ListForUser(ctx, userID)
if err != nil || len(personas) == 0 {
return ""
}
var lines []string
for _, p := range personas {
if !p.IsActive || p.Handle == "" || p.ID == currentPersonaID {
continue
}
desc := p.Description
if len(desc) > 80 {
desc = desc[:80] + "..."
}
if desc != "" {
lines = append(lines, fmt.Sprintf("- @%s (%s) — %s", p.Handle, p.Name, desc))
} else {
lines = append(lines, fmt.Sprintf("- @%s (%s)", p.Handle, p.Name))
}
}
if len(lines) == 0 {
return ""
}
hint := "Available conversation participants you can @mention to direct a response:\n"
for _, l := range lines {
hint += l + "\n"
}
hint += "To ask another participant to respond, include their @handle in your message."
return hint
}
// Resolves @tokens in message content directly against the model catalog
// and personas table. Works in any chat — no channel roster needed.
//
// Resolution order:
// 1. Persona handle (exact match, then prefix)
// 2. User handle/username (exact, then prefix) — v0.23.1
// When a user is @mentioned in a DM/channel, skip completion and deliver
// a notification instead. Returns non-empty mentionedUserID in that case.
// 3. Model ID in enabled catalog (exact match, then prefix)
//
// Returns (modelID, providerConfigID, *Persona, mentionedUserID).
// Exactly one of (modelID, mentionedUserID) will be non-empty on a successful
// match, or all empty if no @mention found/resolved.
func (h *CompletionHandler) resolveMention(ctx context.Context, userID, content string) (string, string, *models.Persona, string) {
// Extract first @token
token := extractFirstMention(content)
if token == "" {
return "", "", nil, ""
}
// Normalize: lowercase, hyphens
normalized := strings.ToLower(strings.TrimSpace(token))
// 1. Try persona handle (exact)
var personaID, personaHandle string
err := database.DB.QueryRowContext(ctx, database.Q(`
SELECT id, handle FROM personas
WHERE LOWER(handle) = $1 AND is_active = true
`), normalized).Scan(&personaID, &personaHandle)
if err == nil && personaID != "" {
p := ResolvePersona(h.stores, personaID, userID)
if p != nil {
cfgID := ""
if p.ProviderConfigID != nil {
cfgID = *p.ProviderConfigID
}
log.Printf("[mention] @%s → persona %s (%s)", token, p.Name, p.BaseModelID)
return p.BaseModelID, cfgID, p, ""
}
}
// 2. Try persona handle (prefix — unambiguous only)
if personaID == "" {
var count int
database.DB.QueryRowContext(ctx, database.Q(`
SELECT COUNT(*) FROM personas WHERE LOWER(handle) LIKE $1 AND is_active = true
`), normalized+"%").Scan(&count)
if count == 1 {
database.DB.QueryRowContext(ctx, database.Q(`
SELECT id FROM personas WHERE LOWER(handle) LIKE $1 AND is_active = true
`), normalized+"%").Scan(&personaID)
if personaID != "" {
p := ResolvePersona(h.stores, personaID, userID)
if p != nil {
cfgID := ""
if p.ProviderConfigID != nil {
cfgID = *p.ProviderConfigID
}
log.Printf("[mention] @%s → persona prefix %s (%s)", token, p.Name, p.BaseModelID)
return p.BaseModelID, cfgID, p, ""
}
}
}
}
// 3. Try user handle (exact match) — v0.24.0
// Only resolves to a user if the caller is not the same user (no self-mention)
var mentionedUserID string
err = database.DB.QueryRowContext(ctx, database.Q(`
SELECT id FROM users
WHERE LOWER(handle) = $1 AND id != $2 AND is_active = true
LIMIT 1
`), normalized, userID).Scan(&mentionedUserID)
if err == nil && mentionedUserID != "" {
log.Printf("[mention] @%s → user %s", token, mentionedUserID)
return "", "", nil, mentionedUserID
}
// 4. Try user handle (prefix — unambiguous only) — v0.24.0
var userCount int
database.DB.QueryRowContext(ctx, database.Q(`
SELECT COUNT(*) FROM users
WHERE LOWER(handle) LIKE $1 AND id != $2 AND is_active = true
`), normalized+"%", userID).Scan(&userCount)
if userCount == 1 {
database.DB.QueryRowContext(ctx, database.Q(`
SELECT id FROM users
WHERE LOWER(handle) LIKE $1 AND id != $2 AND is_active = true
LIMIT 1
`), normalized+"%", userID).Scan(&mentionedUserID)
if mentionedUserID != "" {
log.Printf("[mention] @%s → user prefix %s", token, mentionedUserID)
return "", "", nil, mentionedUserID
}
}
var modelID, provConfigID string
err = database.DB.QueryRowContext(ctx, database.Q(`
SELECT mc.model_id, mc.provider_config_id
FROM model_catalog mc
JOIN provider_configs pc ON pc.id = mc.provider_config_id
WHERE LOWER(mc.model_id) = $1
AND mc.visibility = 'enabled'
AND pc.is_active = true
ORDER BY
CASE pc.scope WHEN 'global' THEN 0 WHEN 'team' THEN 1 WHEN 'personal' THEN 2 END
LIMIT 1
`), normalized).Scan(&modelID, &provConfigID)
if err == nil && modelID != "" {
log.Printf("[mention] @%s → model %s (config %s)", token, modelID, provConfigID)
return modelID, provConfigID, nil, ""
}
// 5. Try model_id prefix (unambiguous)
var prefixCount int
database.DB.QueryRowContext(ctx, database.Q(`
SELECT COUNT(DISTINCT mc.model_id)
FROM model_catalog mc
JOIN provider_configs pc ON pc.id = mc.provider_config_id
WHERE LOWER(mc.model_id) LIKE $1
AND mc.visibility = 'enabled'
AND pc.is_active = true
`), normalized+"%").Scan(&prefixCount)
if prefixCount == 1 {
database.DB.QueryRowContext(ctx, database.Q(`
SELECT mc.model_id, mc.provider_config_id
FROM model_catalog mc
JOIN provider_configs pc ON pc.id = mc.provider_config_id
WHERE LOWER(mc.model_id) LIKE $1
AND mc.visibility = 'enabled'
AND pc.is_active = true
ORDER BY CASE pc.scope WHEN 'global' THEN 0 WHEN 'team' THEN 1 WHEN 'personal' THEN 2 END
LIMIT 1
`), normalized+"%").Scan(&modelID, &provConfigID)
if modelID != "" {
log.Printf("[mention] @%s → model prefix %s (config %s)", token, modelID, provConfigID)
return modelID, provConfigID, nil, ""
}
}
return "", "", nil, ""
}
// extractFirstMention finds the first @token in content.
// Returns the token without the @ prefix, or empty string.
func extractFirstMention(content string) string {
for i := 0; i < len(content); i++ {
if content[i] != '@' {
continue
}
// @ must be at start or preceded by whitespace
if i > 0 && content[i-1] != ' ' && content[i-1] != '\n' && content[i-1] != '\t' {
continue
}
// Extract token
start := i + 1
end := start
for end < len(content) && content[end] != ' ' && content[end] != '\n' && content[end] != '\t' {
end++
}
if end > start {
// Strip trailing punctuation
for end > start && (content[end-1] == ',' || content[end-1] == '.' || content[end-1] == '!' || content[end-1] == '?' || content[end-1] == ':' || content[end-1] == ';') {
end--
}
if end > start {
return content[start:end]
}
}
}
return ""
}
// ── Config Resolution ───────────────────────
// Priority: request.provider_config_id → chat.provider_config_id → user's first active config
func (h *CompletionHandler) resolveConfig(userID string, channelID string, req completionRequest) (providers.ProviderConfig, string, string, string, string, error) {
res, err := ResolveProviderConfig(h.stores, h.vault, userID, channelID, req.ProviderConfigID, req.Model)
if err != nil {
return providers.ProviderConfig{}, "", "", "", "", err
}
return res.Config, res.ProviderID, res.Model, res.ConfigID, res.ProviderScope, nil
}
// ── Conversation Loader ─────────────────────
// loadConversation builds the message history for the LLM by walking the
// active path through the message tree. Only the current branch is sent —
// sibling branches are never included in context.
//
// Summary-aware: if the path contains a summary node (metadata.type = "summary"),
// messages before it are replaced by the summary content as a system message.
func (h *CompletionHandler) loadConversation(channelID, userID, personaSystemPrompt, personaID, currentModel string) ([]providers.Message, error) {
messages := make([]providers.Message, 0)
// ── Admin system prompt (always injected first, no opt out) ──
adminPrompt, err := h.stores.GlobalConfig.Get(context.Background(), "system_prompt")
if err == nil {
if content, ok := adminPrompt["content"].(string); ok && content != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: content,
})
}
}
// ── User/persona system prompt (appended after admin prompt) ──
var systemPrompt *string
_ = database.DB.QueryRow(
database.Q(`SELECT system_prompt FROM channels WHERE id = $1`), channelID,
).Scan(&systemPrompt)
// Persona system prompt takes priority; channel system prompt is fallback
if personaSystemPrompt != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: personaSystemPrompt,
})
} else if systemPrompt != nil && *systemPrompt != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: *systemPrompt,
})
}
// ── Project system prompt (v0.19.1) ──
if h.stores.Projects != nil {
projID, _ := h.stores.Projects.GetProjectIDForChannel(context.Background(), channelID)
if projID != "" {
if proj, err := h.stores.Projects.GetByID(context.Background(), projID); err == nil {
if sp, ok := proj.Settings["system_prompt"].(string); ok && sp != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: sp,
})
}
}
}
}
// ── Pre-completion filter chain (v0.29.0) ──
// Runs registered filters (KB auto-inject, future extension filters).
// Replaces the inline BuildKBHint call. Each filter contributes zero
// or more system messages. Failures are logged and skipped.
if h.filterChain != nil && h.filterChain.Len() > 0 {
cc := &filters.CompletionContext{
ChannelID: channelID,
UserID: userID,
PersonaID: personaID,
}
for _, injected := range h.filterChain.Execute(context.Background(), cc) {
messages = append(messages, providers.Message{
Role: injected.Role,
Content: injected.Content,
})
}
}
// ── Memory hint (inject known facts about this user — v0.18.0) ──
// We pass empty lastUserMessage here; the current message content is available
// at the call site but not in loadConversation. Semantic recall will improve
// when the hint is built closer to the user's message.
if memHint := BuildMemoryHint(context.Background(), h.stores, h.embedder, userID, personaID, ""); memHint != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: memHint,
})
}
// ── Participant roster (v0.23.0) ──
// Inject the list of available @mentionable personas so the LLM
// knows who it can direct responses to. This enables LLM→LLM chaining.
if participantHint := h.buildParticipantHint(channelID, userID, personaID); participantHint != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: participantHint,
})
}
// Walk the active path (root → leaf) instead of loading all messages
path, err := getActivePath(channelID, userID)
if err != nil {
return nil, err
}
// Find the most recent summary node in the path
summaryIdx := -1
for i, m := range path {
if isSummaryMessage(&m) {
summaryIdx = i
}
}
// If we found a summary, inject it as a system message and skip prior messages
startIdx := 0
if summaryIdx >= 0 {
messages = append(messages, providers.Message{
Role: "system",
Content: "Previous conversation summary:\n" + path[summaryIdx].Content,
})
startIdx = summaryIdx + 1
}
// ── Build persona name cache for history rewrite (v0.23.0) ──
// Collect unique persona IDs from the path so we can attribute
// foreign persona messages by name instead of raw UUID.
personaNameCache := make(map[string]string)
hasForeignPersona := false
{
seen := make(map[string]bool)
for _, m := range path {
if m.ParticipantType == "persona" && m.ParticipantID != "" && m.ParticipantID != personaID {
if !seen[m.ParticipantID] {
seen[m.ParticipantID] = true
}
}
}
for pid := range seen {
var name string
database.DB.QueryRow(database.Q(`SELECT name FROM personas WHERE id = $1`), pid).Scan(&name)
if name != "" {
personaNameCache[pid] = name
}
}
}
for _, m := range path[startIdx:] {
if m.Role == "system" {
continue // system prompts handled above
}
// Skip summary nodes that aren't the boundary (shouldn't happen, but defensive)
if isSummaryMessage(&m) {
continue
}
// ── Foreign speaker attribution (v0.23.0) ──
// ALL assistant messages get a Name field identifying who generated them.
// Providers that support "name" (OpenAI, OpenRouter, Venice) use it
// to distinguish speakers. Anthropic rewrites named foreign messages
// to user role with attribution.
msg := providers.Message{
Role: m.Role,
Content: m.Content,
}
if m.Role == "assistant" {
isForeign := false
if m.ParticipantType == "persona" && m.ParticipantID != "" {
if m.ParticipantID != personaID {
// Foreign persona — tag with name
isForeign = true
name := personaNameCache[m.ParticipantID]
if name != "" {
msg.Name = name
}
}
} else if m.Model != nil && *m.Model != "" {
// Raw model — foreign if model differs or we're a persona
if *m.Model != currentModel || personaID != "" {
isForeign = true
msg.Name = *m.Model
}
}
if isForeign {
hasForeignPersona = true
}
}
messages = append(messages, msg)
}
// If foreign persona messages exist, add a system hint explaining
// the multi-participant context (helps all providers)
if hasForeignPersona {
hint := "This is a multi-participant AI conversation. Messages from other AI participants are marked with their name. You are one of several models — respond only as yourself, ignoring other participants' styles and characters."
// Insert after the last system message, before conversation history
insertIdx := 0
for i, m := range messages {
if m.Role == "system" {
insertIdx = i + 1
} else {
break
}
}
messages = append(messages[:insertIdx+1], messages[insertIdx:]...)
messages[insertIdx] = providers.Message{Role: "system", Content: hint}
}
return messages, nil
}
// ── Message Persistence ─────────────────────
// persistMessage inserts a message into the tree, updates the cursor, and
// returns the new message's ID.
//
// Parent resolution:
// - parentOverride (explicit parent for edit/regen operations)
// - cursor's active_leaf_id (normal conversation flow)
// - latest message by created_at (fallback for legacy/missing cursor)
// toolActivityJSON marshals tool activity for persistence. Returns nil if empty.
func toolActivityJSON(activity []map[string]interface{}) json.RawMessage {
if len(activity) == 0 {
return nil
}
b, _ := json.Marshal(activity)
return b
}
func (h *CompletionHandler) persistMessage(channelID, userID, role, content, model string, inputTokens, outputTokens int, parentOverride *string, toolCallsJSON json.RawMessage, providerConfigID string, personaID string) (string, error) {
var tokensUsed *int
if inputTokens > 0 || outputTokens > 0 {
total := inputTokens + outputTokens
tokensUsed = &total
}
// Determine parent
var parentID *string
if parentOverride != nil {
parentID = parentOverride
} else {
parentID, _ = getActiveLeaf(channelID, userID)
}
// Compute sibling_index
siblingIdx := nextSiblingIndex(channelID, parentID)
// Determine participant: persona for assistant messages with personaID, otherwise user/model
participantType := "user"
participantID := userID
if role == "assistant" {
if personaID != "" {
participantType = "persona"
participantID = personaID
} else {
participantType = "model"
participantID = model
}
}
// provider_config_id: nil if empty
var provCfgVal interface{}
if providerConfigID != "" {
provCfgVal = providerConfigID
}
// Prepare tool_calls JSONB (nil → NULL)
var tcVal interface{}
if len(toolCallsJSON) > 0 {
tcVal = string(toolCallsJSON)
}
// Insert with RETURNING id (Postgres) or pre-generated id (SQLite)
var newID string
if database.IsSQLite() {
newID = store.NewID()
_, err := database.DB.Exec(`
INSERT INTO messages (id, channel_id, role, content, model, tokens_used,
tool_calls, parent_id, participant_type, participant_id, provider_config_id, sibling_index)
VALUES (?, ?, ?, ?, NULLIF(?, ''), ?, ?, ?, ?, ?, ?, ?)
`, newID, channelID, role, content, model, tokensUsed,
tcVal, parentID, participantType, participantID, provCfgVal, siblingIdx)
if err != nil {
return "", err
}
} else {
err := database.DB.QueryRow(`
INSERT INTO messages (channel_id, role, content, model, tokens_used,
tool_calls, parent_id, participant_type, participant_id, provider_config_id, sibling_index)
VALUES ($1, $2, $3, NULLIF($4, ''), $5, $6, $7, $8, $9, $10, $11)
RETURNING id
`, channelID, role, content, model, tokensUsed,
tcVal, parentID, participantType, participantID, provCfgVal, siblingIdx,
).Scan(&newID)
if err != nil {
return "", err
}
}
// Update cursor to point at new message
if userID != "" {
_ = updateCursor(channelID, userID, newID)
}
// Touch channel updated_at
_, _ = database.DB.Exec(database.Q(`UPDATE channels SET updated_at = NOW() WHERE id = $1`), channelID)
return newID, nil
}
// ── Usage Logging ─────────────────────────
// logUsage records token usage and cost in the usage_log table.
// Always logs even if tokens are zero — the request happened.
func (h *CompletionHandler) logUsage(
c *gin.Context,
channelID, userID, configID, providerScope, modelID string,
inputTokens, outputTokens, cacheCreation, cacheRead int,
) {
if h.stores.Usage == nil {
return
}
entry := &models.UsageEntry{
ChannelID: &channelID,
UserID: userID,
ProviderConfigID: &configID,
ProviderScope: providerScope,
ModelID: modelID,
PromptTokens: inputTokens,
CompletionTokens: outputTokens,
CacheCreationTokens: cacheCreation,
CacheReadTokens: cacheRead,
}
// Look up pricing for cost calculation
if h.stores.Pricing != nil {
pricing, err := h.stores.Pricing.GetForModel(c.Request.Context(), configID, modelID)
if err == nil && pricing != nil {
entry.CostInput = calcCost(inputTokens, pricing.InputPerM)
entry.CostOutput = calcCost(outputTokens, pricing.OutputPerM)
}
}
if err := h.stores.Usage.Log(c.Request.Context(), entry); err != nil {
log.Printf("⚠ Failed to log usage: %v", err)
}
// v0.33.0: Prometheus token counters
metrics.CompletionTokensTotal.WithLabelValues("prompt", modelID).Add(float64(inputTokens))
metrics.CompletionTokensTotal.WithLabelValues("completion", modelID).Add(float64(outputTokens))
// v0.33.0: Structured usage log for observability correlation
correlationID, _ := c.Get("request_id")
slog.Info("completion.usage",
"request_id", correlationID,
"model_id", modelID,
"provider_config_id", configID,
"input_tokens", inputTokens,
"output_tokens", outputTokens,
)
}
// calcCost computes the cost for a given token count and price-per-million.
func calcCost(tokens int, pricePerM *float64) *float64 {
if pricePerM == nil || tokens == 0 {
return nil
}
cost := float64(tokens) * *pricePerM / 1_000_000.0
return &cost
}
// recordHealth records provider call outcome for health tracking (v0.22.0)
// and Prometheus completion metrics (v0.33.0).
func (h *CompletionHandler) recordHealth(configID string, start time.Time, err error) {
duration := time.Since(start)
latencyMs := int(duration.Milliseconds())
// v0.33.0: Prometheus completion duration (always, even on error)
if configID != "" {
metrics.CompletionDuration.WithLabelValues(configID, "").Observe(duration.Seconds())
}
if h.health == nil || configID == "" {
return
}
if err != nil {
errMsg := err.Error()
status := "error"
if strings.Contains(errMsg, "context deadline exceeded") || strings.Contains(errMsg, "timeout") {
h.health.RecordTimeout(configID, latencyMs, errMsg)
status = "timeout"
} else if strings.Contains(errMsg, "HTTP 429") || strings.Contains(errMsg, "rate limit") || strings.Contains(errMsg, "Too Many Requests") {
h.health.RecordRateLimit(configID, latencyMs, errMsg)
status = "rate_limited"
} else {
h.health.RecordError(configID, latencyMs, errMsg)
}
metrics.CompletionsTotal.WithLabelValues(configID, "", status).Inc()
} else {
h.health.RecordSuccess(configID, latencyMs)
metrics.CompletionsTotal.WithLabelValues(configID, "", "success").Inc()
}
}
// truncateContent returns content truncated to maxLen runes with ellipsis.
func truncateContent(s string, maxLen int) string {
runes := []rune(s)
if len(runes) <= maxLen {
return s
}
return string(runes[:maxLen]) + "…"
}