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
2026-02-19 15:03:20 +00:00

418 lines
12 KiB
Go

package handlers
import (
"database/sql"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/providers"
)
// ── Request Types ───────────────────────────
type completionRequest struct {
ChatID string `json:"chat_id" binding:"required"`
Content string `json:"content" binding:"required"`
Model string `json:"model,omitempty"`
APIConfigID string `json:"api_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"`
}
// CompletionHandler proxies LLM requests through the backend.
type CompletionHandler struct{}
// NewCompletionHandler creates a new handler.
func NewCompletionHandler() *CompletionHandler {
return &CompletionHandler{}
}
// ── 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 chat_messages
// 4. Persist user message
// 5. Call provider (stream or non-stream)
// 6. Stream SSE to client / return JSON
// 7. 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
}
userID := getUserID(c)
// Verify chat ownership
if !userOwnsChat(c, req.ChatID, userID) {
return
}
// Resolve provider config
providerCfg, providerID, model, configID, err := h.resolveConfig(userID, req)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
provider, err := providers.Get(providerID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
// Load conversation history
messages, err := h.loadConversation(req.ChatID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load conversation"})
return
}
// Add the new user message
messages = append(messages, providers.Message{
Role: "user",
Content: req.Content,
})
// Persist user message
if err := h.persistMessage(req.ChatID, "user", req.Content, "", 0, 0); err != nil {
log.Printf("Failed to persist user message: %v", err)
}
// Build provider request
provReq := providers.CompletionRequest{
Model: model,
Messages: messages,
}
// Resolve capabilities for this model — auto-set defaults
caps := h.getModelCapabilities(model, configID)
if req.MaxTokens > 0 {
provReq.MaxTokens = req.MaxTokens
} else {
// ResolveMaxOutput checks: caps → known models → context/8 → 4096
provReq.MaxTokens = providers.ResolveMaxOutput(model, caps)
}
if req.Temperature != nil {
provReq.Temperature = req.Temperature
}
if req.TopP != nil {
provReq.TopP = req.TopP
}
// Determine streaming
stream := true // default
if req.Stream != nil {
stream = *req.Stream
}
if stream {
h.streamCompletion(c, provider, providerCfg, provReq, req.ChatID, model)
} else {
h.syncCompletion(c, provider, providerCfg, provReq, req.ChatID, model)
}
}
// ── Streaming Completion (SSE) ──────────────
func (h *CompletionHandler) streamCompletion(
c *gin.Context,
provider providers.Provider,
cfg providers.ProviderConfig,
req providers.CompletionRequest,
chatID, model string,
) {
ch, err := provider.StreamCompletion(c.Request.Context(), cfg, req)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": "provider error: " + err.Error()})
return
}
// Set SSE headers
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no") // Disable nginx buffering
c.Status(http.StatusOK)
var fullContent string
flusher, _ := c.Writer.(http.Flusher)
for event := range ch {
if event.Error != nil {
fmt.Fprintf(c.Writer, "data: {\"error\":\"%s\"}\n\n", event.Error.Error())
if flusher != nil {
flusher.Flush()
}
break
}
if event.Delta != "" {
fullContent += event.Delta
// OpenAI-compatible SSE format
fmt.Fprintf(c.Writer, "data: {\"choices\":[{\"delta\":{\"content\":%q},\"finish_reason\":null}],\"model\":%q}\n\n",
event.Delta, model)
if flusher != nil {
flusher.Flush()
}
}
if event.Done {
finishReason := event.FinishReason
if finishReason == "" {
finishReason = "stop"
}
fmt.Fprintf(c.Writer, "data: {\"choices\":[{\"delta\":{},\"finish_reason\":%q}],\"model\":%q}\n\n",
finishReason, model)
io.WriteString(c.Writer, "data: [DONE]\n\n")
if flusher != nil {
flusher.Flush()
}
break
}
}
// Persist assistant response
if fullContent != "" {
if err := h.persistMessage(chatID, "assistant", fullContent, model, 0, 0); err != nil {
log.Printf("Failed to persist assistant message: %v", err)
}
}
}
// ── Non-Streaming Completion ────────────────
func (h *CompletionHandler) syncCompletion(
c *gin.Context,
provider providers.Provider,
cfg providers.ProviderConfig,
req providers.CompletionRequest,
chatID, model string,
) {
resp, err := provider.ChatCompletion(c.Request.Context(), cfg, req)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{"error": "provider error: " + err.Error()})
return
}
// Persist assistant response
if err := h.persistMessage(chatID, "assistant", resp.Content, model, resp.InputTokens, resp.OutputTokens); err != nil {
log.Printf("Failed to persist assistant message: %v", err)
}
// Return OpenAI-compatible response
c.JSON(http.StatusOK, gin.H{
"model": resp.Model,
"choices": []gin.H{
{
"message": gin.H{
"role": "assistant",
"content": resp.Content,
},
"finish_reason": resp.FinishReason,
},
},
"usage": gin.H{
"prompt_tokens": resp.InputTokens,
"completion_tokens": resp.OutputTokens,
"total_tokens": resp.InputTokens + resp.OutputTokens,
},
})
}
// ── Model Capabilities ──────────────────────
// getModelCapabilities looks up capabilities from model_configs DB,
// then overlays with known model defaults and heuristic detection.
func (h *CompletionHandler) getModelCapabilities(model, apiConfigID string) providers.ModelCapabilities {
// Start with known table or heuristic
// Start with known model table or heuristics
caps, found := providers.LookupKnownModel(model)
if !found {
caps = providers.InferCapabilities(model)
}
// Overlay with DB-stored capabilities (from provider sync or admin edit — authoritative)
var capsJSON []byte
err := database.DB.QueryRow(`
SELECT capabilities FROM model_configs
WHERE model_id = $1 AND api_config_id = $2
`, model, apiConfigID).Scan(&capsJSON)
if err == nil && capsJSON != nil {
var dbCaps providers.ModelCapabilities
if err := json.Unmarshal(capsJSON, &dbCaps); err == nil {
if dbCaps.HasProviderData() {
// DB has real provider data — use as authoritative, fill gaps
caps = providers.MergeCapabilities(dbCaps, model)
}
}
}
return caps
}
// ── Config Resolution ───────────────────────
// Priority: request.api_config_id → chat.api_config_id → user's first active config
func (h *CompletionHandler) resolveConfig(userID string, req completionRequest) (providers.ProviderConfig, string, string, string, error) {
var configID string
// 1. Explicit config from request
if req.APIConfigID != "" {
configID = req.APIConfigID
}
// 2. Config from chat
if configID == "" {
var chatConfigID *string
err := database.DB.QueryRow(
`SELECT api_config_id FROM chats WHERE id = $1`, req.ChatID,
).Scan(&chatConfigID)
if err == nil && chatConfigID != nil {
configID = *chatConfigID
}
}
// 3. User's first active config (personal first, then global)
if configID == "" {
err := database.DB.QueryRow(`
SELECT id FROM api_configs
WHERE (user_id = $1 OR is_global = true) AND is_active = true
ORDER BY user_id NULLS LAST, created_at ASC
LIMIT 1
`, userID).Scan(&configID)
if err != nil {
return providers.ProviderConfig{}, "", "", "", fmt.Errorf("no API config found — add one at /api-configs")
}
}
// Load the config including custom headers and provider settings
var providerID, endpoint string
var apiKey, modelDefault *string
var customHeadersJSON, providerSettingsJSON []byte
err := database.DB.QueryRow(`
SELECT provider, endpoint, api_key_encrypted, model_default, custom_headers, provider_settings
FROM api_configs
WHERE id = $1 AND (user_id = $2 OR is_global = true) AND is_active = true
`, configID, userID).Scan(&providerID, &endpoint, &apiKey, &modelDefault, &customHeadersJSON, &providerSettingsJSON)
if err == sql.ErrNoRows {
return providers.ProviderConfig{}, "", "", "", fmt.Errorf("API config not found or not accessible")
}
if err != nil {
return providers.ProviderConfig{}, "", "", "", fmt.Errorf("failed to load API config: %w", err)
}
// Resolve model: request > config default
model := req.Model
if model == "" && modelDefault != nil {
model = *modelDefault
}
if model == "" {
return providers.ProviderConfig{}, "", "", "", fmt.Errorf("no model specified and no default model in config")
}
key := ""
if apiKey != nil {
key = *apiKey
}
// Parse custom headers
customHeaders := make(map[string]string)
if customHeadersJSON != nil {
_ = json.Unmarshal(customHeadersJSON, &customHeaders)
}
// Parse provider-specific settings
providerSettings := make(map[string]interface{})
if providerSettingsJSON != nil {
_ = json.Unmarshal(providerSettingsJSON, &providerSettings)
}
return providers.ProviderConfig{
Endpoint: endpoint,
APIKey: key,
CustomHeaders: customHeaders,
Settings: providerSettings,
}, providerID, model, configID, nil
}
// ── Conversation Loader ─────────────────────
func (h *CompletionHandler) loadConversation(chatID string) ([]providers.Message, error) {
// Load system prompt from chat
var systemPrompt *string
_ = database.DB.QueryRow(
`SELECT system_prompt FROM chats WHERE id = $1`, chatID,
).Scan(&systemPrompt)
messages := make([]providers.Message, 0)
if systemPrompt != nil && *systemPrompt != "" {
messages = append(messages, providers.Message{
Role: "system",
Content: *systemPrompt,
})
}
// Load message history (oldest first)
rows, err := database.DB.Query(`
SELECT role, content FROM chat_messages
WHERE chat_id = $1
ORDER BY created_at ASC
`, chatID)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var msg providers.Message
if err := rows.Scan(&msg.Role, &msg.Content); err != nil {
return nil, err
}
messages = append(messages, msg)
}
return messages, nil
}
// ── Message Persistence ─────────────────────
func (h *CompletionHandler) persistMessage(chatID, role, content, model string, inputTokens, outputTokens int) error {
var tokensUsed *int
if inputTokens > 0 || outputTokens > 0 {
total := inputTokens + outputTokens
tokensUsed = &total
}
_, err := database.DB.Exec(`
INSERT INTO chat_messages (chat_id, role, content, model, tokens_used)
VALUES ($1, $2, $3, NULLIF($4, ''), $5)
`, chatID, role, content, model, tokensUsed)
if err != nil {
return err
}
// Touch chat updated_at
_, _ = database.DB.Exec(`UPDATE chats SET updated_at = NOW() WHERE id = $1`, chatID)
return nil
}