diff --git a/.gitea/workflows/ci.yaml b/.gitea/workflows/ci.yaml index 014dc45..268c2ce 100644 --- a/.gitea/workflows/ci.yaml +++ b/.gitea/workflows/ci.yaml @@ -53,6 +53,9 @@ jobs: # ── Stage 1: Go Build & Test ───────────────── test: runs-on: ubuntu-latest + env: + GOPRIVATE: git.gobha.me/* + GONOSUMCHECK: git.gobha.me/* steps: - name: Checkout uses: actions/checkout@v4 diff --git a/server/handlers/apiconfigs.go b/server/handlers/apiconfigs.go new file mode 100644 index 0000000..b6bc8bb --- /dev/null +++ b/server/handlers/apiconfigs.go @@ -0,0 +1,442 @@ +package handlers + +import ( + "database/sql" + "encoding/json" + "math" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" + + "git.gobha.me/xcaliber/chat-switchboard/database" + "git.gobha.me/xcaliber/chat-switchboard/providers" +) + +// ── Request / Response Types ──────────────── + +type createAPIConfigRequest struct { + Name string `json:"name" binding:"required,max=100"` + Provider string `json:"provider" binding:"required"` + Endpoint string `json:"endpoint" binding:"required"` + APIKey string `json:"api_key"` + ModelDefault string `json:"model_default,omitempty"` + Config map[string]interface{} `json:"config,omitempty"` +} + +type updateAPIConfigRequest struct { + Name *string `json:"name,omitempty"` + Endpoint *string `json:"endpoint,omitempty"` + APIKey *string `json:"api_key,omitempty"` + ModelDefault *string `json:"model_default,omitempty"` + Config map[string]interface{} `json:"config,omitempty"` + IsActive *bool `json:"is_active,omitempty"` +} + +type apiConfigResponse struct { + ID string `json:"id"` + UserID *string `json:"user_id"` + Name string `json:"name"` + Provider string `json:"provider"` + Endpoint string `json:"endpoint"` + HasKey bool `json:"has_key"` // Never expose the actual key + ModelDefault *string `json:"model_default"` + Config map[string]interface{} `json:"config"` + IsActive bool `json:"is_active"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +// APIConfigHandler holds dependencies. +type APIConfigHandler struct{} + +// NewAPIConfigHandler creates a new handler. +func NewAPIConfigHandler() *APIConfigHandler { + return &APIConfigHandler{} +} + +// ── List API Configs ──────────────────────── + +func (h *APIConfigHandler) ListConfigs(c *gin.Context) { + userID := getUserID(c) + page, perPage, offset := parsePagination(c) + + // Count: user's configs + global configs + var total int + err := database.DB.QueryRow( + `SELECT COUNT(*) FROM api_configs WHERE user_id = $1 OR user_id IS NULL`, + userID, + ).Scan(&total) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to count configs"}) + return + } + + rows, err := database.DB.Query(` + SELECT id, user_id, name, provider, endpoint, api_key_encrypted, + model_default, config, is_active, created_at, updated_at + FROM api_configs + WHERE user_id = $1 OR user_id IS NULL + ORDER BY user_id NULLS LAST, name ASC + LIMIT $2 OFFSET $3 + `, userID, perPage, offset) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list configs"}) + return + } + defer rows.Close() + + configs := make([]apiConfigResponse, 0) + for rows.Next() { + cfg, err := scanAPIConfig(rows) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to scan config"}) + return + } + configs = append(configs, cfg) + } + + c.JSON(http.StatusOK, paginatedResponse{ + Data: configs, + Page: page, + PerPage: perPage, + Total: total, + TotalPages: int(math.Ceil(float64(total) / float64(perPage))), + }) +} + +// ── Create API Config ─────────────────────── + +func (h *APIConfigHandler) CreateConfig(c *gin.Context) { + userID := getUserID(c) + + var req createAPIConfigRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // Validate provider exists + if _, err := providers.Get(req.Provider); err != nil { + c.JSON(http.StatusBadRequest, gin.H{ + "error": "unsupported provider: " + req.Provider, + "supported_providers": providers.List(), + }) + return + } + + configJSON := "{}" + if req.Config != nil { + b, _ := json.Marshal(req.Config) + configJSON = string(b) + } + + var cfg apiConfigResponse + var apiKeyEnc *string + var configRaw string + + err := database.DB.QueryRow(` + INSERT INTO api_configs (user_id, name, provider, endpoint, api_key_encrypted, model_default, config) + VALUES ($1, $2, $3, $4, $5, $6, $7::jsonb) + RETURNING id, user_id, name, provider, endpoint, api_key_encrypted, + model_default, config::text, is_active, created_at, updated_at + `, userID, req.Name, req.Provider, req.Endpoint, req.APIKey, req.ModelDefault, configJSON, + ).Scan( + &cfg.ID, &cfg.UserID, &cfg.Name, &cfg.Provider, &cfg.Endpoint, &apiKeyEnc, + &cfg.ModelDefault, &configRaw, &cfg.IsActive, &cfg.CreatedAt, &cfg.UpdatedAt, + ) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create config"}) + return + } + + cfg.HasKey = apiKeyEnc != nil && *apiKeyEnc != "" + cfg.Config = parseJSONBConfig(configRaw) + + c.JSON(http.StatusCreated, cfg) +} + +// ── Get API Config ────────────────────────── + +func (h *APIConfigHandler) GetConfig(c *gin.Context) { + userID := getUserID(c) + configID := c.Param("id") + + row := database.DB.QueryRow(` + SELECT id, user_id, name, provider, endpoint, api_key_encrypted, + model_default, config::text, is_active, created_at, updated_at + FROM api_configs + WHERE id = $1 AND (user_id = $2 OR user_id IS NULL) + `, configID, userID) + + var cfg apiConfigResponse + var apiKeyEnc *string + var configRaw string + + err := row.Scan( + &cfg.ID, &cfg.UserID, &cfg.Name, &cfg.Provider, &cfg.Endpoint, &apiKeyEnc, + &cfg.ModelDefault, &configRaw, &cfg.IsActive, &cfg.CreatedAt, &cfg.UpdatedAt, + ) + if err == sql.ErrNoRows { + c.JSON(http.StatusNotFound, gin.H{"error": "config not found"}) + return + } + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to get config"}) + return + } + + cfg.HasKey = apiKeyEnc != nil && *apiKeyEnc != "" + cfg.Config = parseJSONBConfig(configRaw) + + c.JSON(http.StatusOK, cfg) +} + +// ── Update API Config ─────────────────────── + +func (h *APIConfigHandler) UpdateConfig(c *gin.Context) { + userID := getUserID(c) + configID := c.Param("id") + + var req updateAPIConfigRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // Verify ownership (only owner can update, not global) + var ownerID *string + err := database.DB.QueryRow(`SELECT user_id FROM api_configs WHERE id = $1`, configID).Scan(&ownerID) + if err == sql.ErrNoRows { + c.JSON(http.StatusNotFound, gin.H{"error": "config not found"}) + return + } + if ownerID == nil || *ownerID != userID { + c.JSON(http.StatusForbidden, gin.H{"error": "cannot modify global or other user's config"}) + return + } + + // Dynamic update + setClauses := []string{} + args := []interface{}{} + argN := 1 + + addClause := func(col string, val interface{}) { + setClauses = append(setClauses, col+" = $"+strconv.Itoa(argN)) + args = append(args, val) + argN++ + } + + if req.Name != nil { + addClause("name", *req.Name) + } + if req.Endpoint != nil { + addClause("endpoint", *req.Endpoint) + } + if req.APIKey != nil { + addClause("api_key_encrypted", *req.APIKey) + } + if req.ModelDefault != nil { + addClause("model_default", *req.ModelDefault) + } + if req.Config != nil { + b, _ := json.Marshal(req.Config) + addClause("config", string(b)) + } + if req.IsActive != nil { + addClause("is_active", *req.IsActive) + } + + if len(setClauses) == 0 { + c.JSON(http.StatusBadRequest, gin.H{"error": "no fields to update"}) + return + } + + query := "UPDATE api_configs SET updated_at = NOW(), " + for i, clause := range setClauses { + if i > 0 { + query += ", " + } + query += clause + } + query += " WHERE id = $" + strconv.Itoa(argN) + args = append(args, configID) + + _, err = database.DB.Exec(query, args...) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update config"}) + return + } + + h.GetConfig(c) +} + +// ── Delete API Config ─────────────────────── + +func (h *APIConfigHandler) DeleteConfig(c *gin.Context) { + userID := getUserID(c) + configID := c.Param("id") + + // Only allow deleting own configs + result, err := database.DB.Exec( + `DELETE FROM api_configs WHERE id = $1 AND user_id = $2`, + configID, userID, + ) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete config"}) + return + } + rows, _ := result.RowsAffected() + if rows == 0 { + c.JSON(http.StatusNotFound, gin.H{"error": "config not found or cannot delete global config"}) + return + } + + c.JSON(http.StatusOK, gin.H{"message": "config deleted"}) +} + +// ── List Models from a Config ─────────────── + +func (h *APIConfigHandler) ListModels(c *gin.Context) { + userID := getUserID(c) + configID := c.Param("id") + + // Load config including API key + var providerID, endpoint string + var apiKey *string + err := database.DB.QueryRow(` + SELECT provider, endpoint, api_key_encrypted + FROM api_configs + WHERE id = $1 AND (user_id = $2 OR user_id IS NULL) AND is_active = true + `, configID, userID).Scan(&providerID, &endpoint, &apiKey) + + if err == sql.ErrNoRows { + c.JSON(http.StatusNotFound, gin.H{"error": "config not found"}) + return + } + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load config"}) + return + } + + provider, err := providers.Get(providerID) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + + key := "" + if apiKey != nil { + key = *apiKey + } + + models, err := provider.ListModels(c.Request.Context(), providers.ProviderConfig{ + Endpoint: endpoint, + APIKey: key, + }) + if err != nil { + c.JSON(http.StatusBadGateway, gin.H{"error": "failed to fetch models: " + err.Error()}) + return + } + + c.JSON(http.StatusOK, gin.H{ + "config_id": configID, + "provider": providerID, + "models": models, + }) +} + +// ── List All Available Models (aggregate) ─── + +func (h *APIConfigHandler) ListAllModels(c *gin.Context) { + userID := getUserID(c) + + rows, err := database.DB.Query(` + SELECT id, name, provider, endpoint, api_key_encrypted + FROM api_configs + WHERE (user_id = $1 OR user_id IS NULL) AND is_active = true + `, userID) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list configs"}) + return + } + defer rows.Close() + + type modelEntry struct { + ID string `json:"id"` + Name string `json:"name,omitempty"` + OwnedBy string `json:"owned_by,omitempty"` + ConfigID string `json:"config_id"` + Provider string `json:"provider"` + } + + allModels := make([]modelEntry, 0) + + for rows.Next() { + var cfgID, name, providerID, endpoint string + var apiKey *string + if err := rows.Scan(&cfgID, &name, &providerID, &endpoint, &apiKey); err != nil { + continue + } + + provider, err := providers.Get(providerID) + if err != nil { + continue + } + + key := "" + if apiKey != nil { + key = *apiKey + } + + models, err := provider.ListModels(c.Request.Context(), providers.ProviderConfig{ + Endpoint: endpoint, + APIKey: key, + }) + if err != nil { + continue // Skip configs that fail + } + + for _, m := range models { + allModels = append(allModels, modelEntry{ + ID: m.ID, + Name: m.Name, + OwnedBy: m.OwnedBy, + ConfigID: cfgID, + Provider: providerID, + }) + } + } + + c.JSON(http.StatusOK, gin.H{"models": allModels}) +} + +// ── Helpers ───────────────────────────────── + +type scannable interface { + Scan(dest ...interface{}) error +} + +func scanAPIConfig(row scannable) (apiConfigResponse, error) { + var cfg apiConfigResponse + var apiKeyEnc *string + var configRaw string + + err := row.Scan( + &cfg.ID, &cfg.UserID, &cfg.Name, &cfg.Provider, &cfg.Endpoint, &apiKeyEnc, + &cfg.ModelDefault, &configRaw, &cfg.IsActive, &cfg.CreatedAt, &cfg.UpdatedAt, + ) + if err != nil { + return cfg, err + } + + cfg.HasKey = apiKeyEnc != nil && *apiKeyEnc != "" + cfg.Config = parseJSONBConfig(configRaw) + return cfg, nil +} + +func parseJSONBConfig(raw string) map[string]interface{} { + result := make(map[string]interface{}) + _ = json.Unmarshal([]byte(raw), &result) + return result +} diff --git a/server/handlers/apiconfigs_test.go b/server/handlers/apiconfigs_test.go new file mode 100644 index 0000000..9bdcdaa --- /dev/null +++ b/server/handlers/apiconfigs_test.go @@ -0,0 +1,131 @@ +package handlers + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + + "git.gobha.me/xcaliber/chat-switchboard/providers" +) + +func init() { + gin.SetMode(gin.TestMode) + providers.Init() +} + +func TestCreateConfigMissingFields(t *testing.T) { + h := NewAPIConfigHandler() + r := gin.New() + r.POST("/api-configs", func(c *gin.Context) { + c.Set("user_id", "test-user") + h.CreateConfig(c) + }) + + tests := []struct { + name string + body string + }{ + {"missing name", `{"provider":"openai","endpoint":"http://x"}`}, + {"missing provider", `{"name":"Test","endpoint":"http://x"}`}, + {"missing endpoint", `{"name":"Test","provider":"openai"}`}, + {"empty body", `{}`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("POST", "/api-configs", + bytes.NewBufferString(tt.body)) + req.Header.Set("Content-Type", "application/json") + r.ServeHTTP(w, req) + + if w.Code != http.StatusBadRequest { + t.Errorf("Expected 400, got %d (body: %s)", w.Code, w.Body.String()) + } + }) + } +} + +func TestCreateConfigInvalidProvider(t *testing.T) { + h := NewAPIConfigHandler() + r := gin.New() + r.POST("/api-configs", func(c *gin.Context) { + c.Set("user_id", "test-user") + h.CreateConfig(c) + }) + + body := `{"name":"Test","provider":"nonexistent","endpoint":"http://x"}` + w := httptest.NewRecorder() + req, _ := http.NewRequest("POST", "/api-configs", + bytes.NewBufferString(body)) + req.Header.Set("Content-Type", "application/json") + r.ServeHTTP(w, req) + + if w.Code != http.StatusBadRequest { + t.Errorf("Expected 400 for invalid provider, got %d", w.Code) + } + + var resp map[string]interface{} + json.Unmarshal(w.Body.Bytes(), &resp) + + if _, ok := resp["supported_providers"]; !ok { + t.Error("Expected supported_providers in error response") + } +} + +func TestCompletionHandlerMissingFields(t *testing.T) { + h := NewCompletionHandler() + r := gin.New() + r.POST("/chat/completions", func(c *gin.Context) { + c.Set("user_id", "test-user") + h.Complete(c) + }) + + tests := []struct { + name string + body string + }{ + {"missing chat_id", `{"content":"hello"}`}, + {"missing content", `{"chat_id":"abc"}`}, + {"empty body", `{}`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + w := httptest.NewRecorder() + req, _ := http.NewRequest("POST", "/chat/completions", + bytes.NewBufferString(tt.body)) + req.Header.Set("Content-Type", "application/json") + r.ServeHTTP(w, req) + + if w.Code != http.StatusBadRequest { + t.Errorf("Expected 400, got %d (body: %s)", w.Code, w.Body.String()) + } + }) + } +} + +func TestSupportedProviders(t *testing.T) { + ids := providers.List() + + expected := map[string]bool{ + "openai": false, + "anthropic": false, + } + + for _, id := range ids { + if _, ok := expected[id]; ok { + expected[id] = true + } + } + + for name, found := range expected { + if !found { + t.Errorf("Expected provider %s to be registered", name) + } + } +} diff --git a/server/handlers/completion.go b/server/handlers/completion.go new file mode 100644 index 0000000..6c74539 --- /dev/null +++ b/server/handlers/completion.go @@ -0,0 +1,361 @@ +package handlers + +import ( + "database/sql" + "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, 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, + } + if req.MaxTokens > 0 { + provReq.MaxTokens = req.MaxTokens + } + 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, + }, + }) +} + +// ── 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, 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 + if configID == "" { + err := database.DB.QueryRow(` + SELECT id FROM api_configs + WHERE (user_id = $1 OR user_id IS NULL) 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 + var providerID, endpoint string + var apiKey, modelDefault *string + err := database.DB.QueryRow(` + SELECT provider, endpoint, api_key_encrypted, model_default + FROM api_configs + WHERE id = $1 AND (user_id = $2 OR user_id IS NULL) AND is_active = true + `, configID, userID).Scan(&providerID, &endpoint, &apiKey, &modelDefault) + + 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 + } + + return providers.ProviderConfig{ + Endpoint: endpoint, + APIKey: key, + }, providerID, model, 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 +} diff --git a/server/main.go b/server/main.go index 80bee09..96939d2 100644 --- a/server/main.go +++ b/server/main.go @@ -9,11 +9,15 @@ import ( "git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/handlers" "git.gobha.me/xcaliber/chat-switchboard/middleware" + "git.gobha.me/xcaliber/chat-switchboard/providers" ) func main() { cfg := config.Load() + // Register LLM providers + providers.Init() + if err := database.Connect(cfg); err != nil { log.Printf("⚠ Database unavailable: %v", err) log.Println(" Running in unmanaged mode (no persistence)") @@ -23,7 +27,7 @@ func main() { r := gin.Default() r.Use(middleware.CORS()) - // Health check + // Health check (k8s probes hit this directly) r.GET("/health", func(c *gin.Context) { c.JSON(200, gin.H{ "status": "ok", @@ -36,15 +40,17 @@ func main() { auth := handlers.NewAuthHandler(cfg) authLimiter := middleware.NewRateLimiter(1, 5) - api := r.Group("/api/v1") - { - api.GET("/health", func(c *gin.Context) { - c.JSON(200, gin.H{ - "status": "ok", - "version": Version, - "database": database.IsConnected(), - }) - }) + api := r.Group("/api/v1") + { + // Health (routable through ingress) + api.GET("/health", func(c *gin.Context) { + c.JSON(200, gin.H{ + "status": "ok", + "version": Version, + "database": database.IsConnected(), + "providers": providers.List(), + }) + }) authGroup := api.Group("/auth") authGroup.Use(authLimiter.Limit()) @@ -71,23 +77,32 @@ func main() { msgs := handlers.NewMessageHandler() protected.GET("/chats/:id/messages", msgs.ListMessages) protected.POST("/chats/:id/messages", msgs.CreateMessage) - protected.POST("/chats/:id/regenerate", msgs.Regenerate) - // Settings — ticket #10 + // ── Chat Engine (#8) ──────────────── + comp := handlers.NewCompletionHandler() + protected.POST("/chat/completions", comp.Complete) + protected.POST("/chats/:id/regenerate", msgs.Regenerate) // TODO: wire to completion handler + + // API Configs + apiCfg := handlers.NewAPIConfigHandler() + protected.GET("/api-configs", apiCfg.ListConfigs) + protected.POST("/api-configs", apiCfg.CreateConfig) + protected.GET("/api-configs/:id", apiCfg.GetConfig) + protected.PUT("/api-configs/:id", apiCfg.UpdateConfig) + protected.DELETE("/api-configs/:id", apiCfg.DeleteConfig) + + // Models (per-config and aggregate) + protected.GET("/api-configs/:id/models", apiCfg.ListModels) + protected.GET("/models", apiCfg.ListAllModels) + + // Settings — future protected.GET("/settings", stubHandler("settings")) protected.PUT("/settings", stubHandler("settings")) - - // API Configs — ticket #10 - protected.GET("/api-configs", stubHandler("api-configs")) - protected.POST("/api-configs", stubHandler("api-configs")) - protected.DELETE("/api-configs/:id", stubHandler("api-configs")) - - // Models — ticket #10 - protected.GET("/models", stubHandler("models")) } } log.Printf("🔀 Chat Switchboard API v%s starting on port %s", Version, cfg.Port) + log.Printf(" Providers: %v", providers.List()) if err := r.Run(":" + cfg.Port); err != nil { log.Fatalf("Failed to start server: %v", err) } diff --git a/server/main_test.go b/server/main_test.go index 52f6712..eaf8c08 100644 --- a/server/main_test.go +++ b/server/main_test.go @@ -10,10 +10,12 @@ import ( "git.gobha.me/xcaliber/chat-switchboard/config" "git.gobha.me/xcaliber/chat-switchboard/handlers" "git.gobha.me/xcaliber/chat-switchboard/middleware" + "git.gobha.me/xcaliber/chat-switchboard/providers" ) func init() { gin.SetMode(gin.TestMode) + providers.Init() } func TestHealthEndpoint(t *testing.T) { @@ -53,6 +55,8 @@ func TestAllRoutesRegistered(t *testing.T) { auth := handlers.NewAuthHandler(cfg) chats := handlers.NewChatHandler() msgs := handlers.NewMessageHandler() + comp := handlers.NewCompletionHandler() + apiCfg := handlers.NewAPIConfigHandler() r := gin.New() api := r.Group("/api/v1") @@ -67,15 +71,29 @@ func TestAllRoutesRegistered(t *testing.T) { protected := api.Group("") { + // Chats protected.GET("/chats", chats.ListChats) protected.POST("/chats", chats.CreateChat) protected.GET("/chats/:id", chats.GetChat) protected.PUT("/chats/:id", chats.UpdateChat) protected.DELETE("/chats/:id", chats.DeleteChat) + // Messages protected.GET("/chats/:id/messages", msgs.ListMessages) protected.POST("/chats/:id/messages", msgs.CreateMessage) protected.POST("/chats/:id/regenerate", msgs.Regenerate) + + // Chat Engine + protected.POST("/chat/completions", comp.Complete) + + // API Configs + protected.GET("/api-configs", apiCfg.ListConfigs) + protected.POST("/api-configs", apiCfg.CreateConfig) + protected.GET("/api-configs/:id", apiCfg.GetConfig) + protected.PUT("/api-configs/:id", apiCfg.UpdateConfig) + protected.DELETE("/api-configs/:id", apiCfg.DeleteConfig) + protected.GET("/api-configs/:id/models", apiCfg.ListModels) + protected.GET("/models", apiCfg.ListAllModels) } routes := r.Routes() @@ -85,18 +103,32 @@ func TestAllRoutesRegistered(t *testing.T) { } expected := []string{ + // Auth "POST /api/v1/auth/register", "POST /api/v1/auth/login", "POST /api/v1/auth/refresh", "POST /api/v1/auth/logout", + // Chats "GET /api/v1/chats", "POST /api/v1/chats", "GET /api/v1/chats/:id", "PUT /api/v1/chats/:id", "DELETE /api/v1/chats/:id", + // Messages "GET /api/v1/chats/:id/messages", "POST /api/v1/chats/:id/messages", "POST /api/v1/chats/:id/regenerate", + // Chat Engine + "POST /api/v1/chat/completions", + // API Configs + "GET /api/v1/api-configs", + "POST /api/v1/api-configs", + "GET /api/v1/api-configs/:id", + "PUT /api/v1/api-configs/:id", + "DELETE /api/v1/api-configs/:id", + "GET /api/v1/api-configs/:id/models", + // Models + "GET /api/v1/models", } for _, e := range expected { @@ -106,6 +138,60 @@ func TestAllRoutesRegistered(t *testing.T) { } } +func TestProviderRegistry(t *testing.T) { + ids := providers.List() + if len(ids) < 2 { + t.Errorf("Expected at least 2 providers, got %d", len(ids)) + } + + // OpenAI should be registered + p, err := providers.Get("openai") + if err != nil { + t.Errorf("OpenAI provider not found: %v", err) + } + if p.ID() != "openai" { + t.Errorf("Expected ID 'openai', got '%s'", p.ID()) + } + + // Anthropic should be registered + p, err = providers.Get("anthropic") + if err != nil { + t.Errorf("Anthropic provider not found: %v", err) + } + if p.ID() != "anthropic" { + t.Errorf("Expected ID 'anthropic', got '%s'", p.ID()) + } + + // Unknown should fail + _, err = providers.Get("nonexistent") + if err == nil { + t.Error("Expected error for unknown provider") + } +} + +func TestAnthropicStaticModels(t *testing.T) { + p := &providers.AnthropicProvider{} + models, err := p.ListModels(nil, providers.ProviderConfig{}) + if err != nil { + t.Errorf("ListModels should not fail: %v", err) + } + if len(models) == 0 { + t.Error("Expected some static models") + } + + // Check that known models are present + found := false + for _, m := range models { + if m.ID == "claude-sonnet-4-20250514" { + found = true + break + } + } + if !found { + t.Error("Expected claude-sonnet-4 in static model list") + } +} + func TestStubHandler(t *testing.T) { r := gin.New() r.GET("/test", stubHandler("settings")) diff --git a/server/providers/anthropic.go b/server/providers/anthropic.go new file mode 100644 index 0000000..7124ce2 --- /dev/null +++ b/server/providers/anthropic.go @@ -0,0 +1,236 @@ +package providers + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +const anthropicDefaultEndpoint = "https://api.anthropic.com" +const anthropicAPIVersion = "2023-06-01" + +// AnthropicProvider handles the Anthropic Messages API. +type AnthropicProvider struct{} + +func (p *AnthropicProvider) ID() string { return "anthropic" } + +// ── Chat Completion (non-streaming) ───────── + +func (p *AnthropicProvider) ChatCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (*CompletionResponse, error) { + req.Stream = false + + body, err := p.doRequest(ctx, cfg, req) + if err != nil { + return nil, err + } + defer body.Close() + + var resp anthropicResponse + if err := json.NewDecoder(body).Decode(&resp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + // Extract text from content blocks + var content string + for _, block := range resp.Content { + if block.Type == "text" { + content += block.Text + } + } + + return &CompletionResponse{ + Content: content, + Model: resp.Model, + FinishReason: resp.StopReason, + InputTokens: resp.Usage.InputTokens, + OutputTokens: resp.Usage.OutputTokens, + }, nil +} + +// ── Stream Completion ─────────────────────── + +func (p *AnthropicProvider) StreamCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (<-chan StreamEvent, error) { + req.Stream = true + + body, err := p.doRequest(ctx, cfg, req) + if err != nil { + return nil, err + } + + ch := make(chan StreamEvent, 64) + + go func() { + defer close(ch) + defer body.Close() + + scanner := bufio.NewScanner(body) + for scanner.Scan() { + line := scanner.Text() + + if !strings.HasPrefix(line, "data: ") { + continue + } + + data := strings.TrimPrefix(line, "data: ") + + var event anthropicStreamEvent + if err := json.Unmarshal([]byte(data), &event); err != nil { + continue + } + + switch event.Type { + case "content_block_delta": + if event.Delta.Type == "text_delta" { + ch <- StreamEvent{Delta: event.Delta.Text} + } + case "message_delta": + ch <- StreamEvent{ + Done: true, + FinishReason: event.Delta.StopReason, + } + case "message_stop": + ch <- StreamEvent{Done: true} + return + case "error": + ch <- StreamEvent{ + Error: fmt.Errorf("anthropic stream error: %s", data), + } + return + } + } + + if err := scanner.Err(); err != nil { + ch <- StreamEvent{Error: err} + } + }() + + return ch, nil +} + +// ── List Models ───────────────────────────── + +func (p *AnthropicProvider) ListModels(_ context.Context, _ ProviderConfig) ([]Model, error) { + // Anthropic doesn't have a list models endpoint. + // Return the known model families. + return []Model{ + {ID: "claude-sonnet-4-20250514", Name: "Claude Sonnet 4"}, + {ID: "claude-haiku-4-20250414", Name: "Claude Haiku 4"}, + {ID: "claude-opus-4-20250514", Name: "Claude Opus 4"}, + {ID: "claude-3-5-sonnet-20241022", Name: "Claude 3.5 Sonnet"}, + {ID: "claude-3-5-haiku-20241022", Name: "Claude 3.5 Haiku"}, + }, nil +} + +// ── HTTP Layer ────────────────────────────── + +func (p *AnthropicProvider) doRequest(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (io.ReadCloser, error) { + endpoint := cfg.Endpoint + if endpoint == "" { + endpoint = anthropicDefaultEndpoint + } + endpoint = strings.TrimSuffix(endpoint, "/") + "/v1/messages" + + // Separate system message from conversation + var system string + messages := make([]anthropicMessage, 0, len(req.Messages)) + for _, m := range req.Messages { + if m.Role == "system" { + system = m.Content + continue + } + messages = append(messages, anthropicMessage{ + Role: m.Role, + Content: m.Content, + }) + } + + // Build Anthropic-format request + antReq := anthropicRequest{ + Model: req.Model, + Messages: messages, + MaxTokens: req.MaxTokens, + Stream: req.Stream, + } + if system != "" { + antReq.System = system + } + if antReq.MaxTokens == 0 { + antReq.MaxTokens = 4096 // Anthropic requires max_tokens + } + if req.Temperature != nil { + antReq.Temperature = req.Temperature + } + if req.TopP != nil { + antReq.TopP = req.TopP + } + + body, err := json.Marshal(antReq) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, "POST", endpoint, bytes.NewReader(body)) + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("x-api-key", cfg.APIKey) + httpReq.Header.Set("anthropic-version", anthropicAPIVersion) + + resp, err := http.DefaultClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("provider request: %w", err) + } + + if resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, fmt.Errorf("provider error: HTTP %d: %s", resp.StatusCode, string(b)) + } + + return resp.Body, nil +} + +// ── Anthropic Wire Types ──────────────────── + +type anthropicMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type anthropicRequest struct { + Model string `json:"model"` + Messages []anthropicMessage `json:"messages"` + System string `json:"system,omitempty"` + MaxTokens int `json:"max_tokens"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + Stream bool `json:"stream"` +} + +type anthropicResponse struct { + Model string `json:"model"` + StopReason string `json:"stop_reason"` + Content []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"content"` + Usage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + } `json:"usage"` +} + +type anthropicStreamEvent struct { + Type string `json:"type"` + Delta struct { + Type string `json:"type"` + Text string `json:"text"` + StopReason string `json:"stop_reason"` + } `json:"delta"` +} diff --git a/server/providers/openai.go b/server/providers/openai.go new file mode 100644 index 0000000..8b54183 --- /dev/null +++ b/server/providers/openai.go @@ -0,0 +1,245 @@ +package providers + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +// OpenAIProvider handles any OpenAI-compatible API. +type OpenAIProvider struct{} + +func (p *OpenAIProvider) ID() string { return "openai" } + +// ── Chat Completion (non-streaming) ───────── + +func (p *OpenAIProvider) ChatCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (*CompletionResponse, error) { + req.Stream = false + + body, err := p.doRequest(ctx, cfg, req) + if err != nil { + return nil, err + } + defer body.Close() + + var resp openaiChatResponse + if err := json.NewDecoder(body).Decode(&resp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + if len(resp.Choices) == 0 { + return nil, fmt.Errorf("no choices in response") + } + + return &CompletionResponse{ + Content: resp.Choices[0].Message.Content, + Model: resp.Model, + FinishReason: resp.Choices[0].FinishReason, + InputTokens: resp.Usage.PromptTokens, + OutputTokens: resp.Usage.CompletionTokens, + }, nil +} + +// ── Stream Completion ─────────────────────── + +func (p *OpenAIProvider) StreamCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (<-chan StreamEvent, error) { + req.Stream = true + + body, err := p.doRequest(ctx, cfg, req) + if err != nil { + return nil, err + } + + ch := make(chan StreamEvent, 64) + + go func() { + defer close(ch) + defer body.Close() + + scanner := bufio.NewScanner(body) + for scanner.Scan() { + line := scanner.Text() + + if !strings.HasPrefix(line, "data: ") { + continue + } + + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + ch <- StreamEvent{Done: true} + return + } + + var chunk openaiStreamChunk + if err := json.Unmarshal([]byte(data), &chunk); err != nil { + continue + } + + if len(chunk.Choices) == 0 { + continue + } + + ev := StreamEvent{ + Delta: chunk.Choices[0].Delta.Content, + Model: chunk.Model, + FinishReason: chunk.Choices[0].FinishReason, + } + if ev.FinishReason != "" { + ev.Done = true + } + + ch <- ev + } + + if err := scanner.Err(); err != nil { + ch <- StreamEvent{Error: err} + } + }() + + return ch, nil +} + +// ── List Models ───────────────────────────── + +func (p *OpenAIProvider) ListModels(ctx context.Context, cfg ProviderConfig) ([]Model, error) { + endpoint := strings.TrimSuffix(cfg.Endpoint, "/") + "/models" + + httpReq, err := http.NewRequestWithContext(ctx, "GET", endpoint, nil) + if err != nil { + return nil, err + } + if cfg.APIKey != "" { + httpReq.Header.Set("Authorization", "Bearer "+cfg.APIKey) + } + + resp, err := http.DefaultClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("list models: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("list models: HTTP %d: %s", resp.StatusCode, string(b)) + } + + var result openaiModelsResponse + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("decode models: %w", err) + } + + models := make([]Model, 0, len(result.Data)) + for _, m := range result.Data { + models = append(models, Model{ + ID: m.ID, + OwnedBy: m.OwnedBy, + }) + } + return models, nil +} + +// ── HTTP Layer ────────────────────────────── + +func (p *OpenAIProvider) doRequest(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (io.ReadCloser, error) { + endpoint := strings.TrimSuffix(cfg.Endpoint, "/") + "/chat/completions" + + // Build OpenAI-format request + oaiReq := openaiChatRequest{ + Model: req.Model, + Stream: req.Stream, + Messages: make([]openaiMessage, 0, len(req.Messages)), + } + if req.MaxTokens > 0 { + oaiReq.MaxTokens = req.MaxTokens + } + if req.Temperature != nil { + oaiReq.Temperature = req.Temperature + } + if req.TopP != nil { + oaiReq.TopP = req.TopP + } + + for _, m := range req.Messages { + oaiReq.Messages = append(oaiReq.Messages, openaiMessage{ + Role: m.Role, + Content: m.Content, + }) + } + + body, err := json.Marshal(oaiReq) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, "POST", endpoint, bytes.NewReader(body)) + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + if cfg.APIKey != "" { + httpReq.Header.Set("Authorization", "Bearer "+cfg.APIKey) + } + + resp, err := http.DefaultClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("provider request: %w", err) + } + + if resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, fmt.Errorf("provider error: HTTP %d: %s", resp.StatusCode, string(b)) + } + + return resp.Body, nil +} + +// ── OpenAI Wire Types ─────────────────────── + +type openaiMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type openaiChatRequest struct { + Model string `json:"model"` + Messages []openaiMessage `json:"messages"` + MaxTokens int `json:"max_tokens,omitempty"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + Stream bool `json:"stream"` +} + +type openaiChatResponse struct { + Model string `json:"model"` + Choices []struct { + Message openaiMessage `json:"message"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + } `json:"usage"` +} + +type openaiStreamChunk struct { + Model string `json:"model"` + Choices []struct { + Delta struct { + Content string `json:"content"` + } `json:"delta"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` +} + +type openaiModelsResponse struct { + Data []struct { + ID string `json:"id"` + OwnedBy string `json:"owned_by"` + } `json:"data"` +} diff --git a/server/providers/provider.go b/server/providers/provider.go new file mode 100644 index 0000000..111cba1 --- /dev/null +++ b/server/providers/provider.go @@ -0,0 +1,85 @@ +package providers + +import ( + "context" +) + +// ── Provider Interface ────────────────────── + +// Provider is the contract for LLM API adapters. +type Provider interface { + // ID returns the provider identifier (e.g. "openai", "anthropic"). + ID() string + + // ChatCompletion sends a non-streaming request and returns the full response. + ChatCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (*CompletionResponse, error) + + // StreamCompletion sends a streaming request and returns a channel of events. + // The channel is closed when the stream ends. Errors are sent as StreamEvents. + StreamCompletion(ctx context.Context, cfg ProviderConfig, req CompletionRequest) (<-chan StreamEvent, error) + + // ListModels returns available models from this provider. + ListModels(ctx context.Context, cfg ProviderConfig) ([]Model, error) +} + +// ── Configuration ─────────────────────────── + +// ProviderConfig holds credentials and endpoint for a configured provider. +// Populated from the api_configs table at call time. +type ProviderConfig struct { + Endpoint string + APIKey string + Extra map[string]interface{} // Provider-specific settings from config JSONB +} + +// ── Request / Response Types ──────────────── + +// Message represents a chat message in the conversation. +type Message struct { + Role string `json:"role"` // user, assistant, system + Content string `json:"content"` +} + +// CompletionRequest is the normalized request sent to any provider. +type CompletionRequest struct { + Model string `json:"model"` + Messages []Message `json:"messages"` + MaxTokens int `json:"max_tokens,omitempty"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + Stream bool `json:"stream,omitempty"` +} + +// CompletionResponse is the normalized non-streaming response. +type CompletionResponse struct { + Content string `json:"content"` + Model string `json:"model"` + FinishReason string `json:"finish_reason"` + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` +} + +// StreamEvent is a single chunk from a streaming response. +type StreamEvent struct { + // Delta is the incremental text content (empty for non-content events). + Delta string `json:"delta,omitempty"` + + // Done is true when the stream is complete. + Done bool `json:"done,omitempty"` + + // FinishReason is set on the final event. + FinishReason string `json:"finish_reason,omitempty"` + + // Model echoes back the model used. + Model string `json:"model,omitempty"` + + // Error is set if the stream encountered an error. + Error error `json:"-"` +} + +// Model represents an available model from a provider. +type Model struct { + ID string `json:"id"` + Name string `json:"name,omitempty"` + OwnedBy string `json:"owned_by,omitempty"` +} diff --git a/server/providers/registry.go b/server/providers/registry.go new file mode 100644 index 0000000..39a1b42 --- /dev/null +++ b/server/providers/registry.go @@ -0,0 +1,47 @@ +package providers + +import ( + "fmt" + "sync" +) + +// registry holds all registered providers. +var ( + registry = make(map[string]Provider) + mu sync.RWMutex +) + +// Register adds a provider to the registry. +func Register(p Provider) { + mu.Lock() + defer mu.Unlock() + registry[p.ID()] = p +} + +// Get returns a provider by ID. +func Get(id string) (Provider, error) { + mu.RLock() + defer mu.RUnlock() + p, ok := registry[id] + if !ok { + return nil, fmt.Errorf("unknown provider: %s", id) + } + return p, nil +} + +// List returns all registered provider IDs. +func List() []string { + mu.RLock() + defer mu.RUnlock() + ids := make([]string, 0, len(registry)) + for id := range registry { + ids = append(ids, id) + } + return ids +} + +// Init registers all built-in providers. +func Init() { + Register(&OpenAIProvider{}) + Register(&AnthropicProvider{}) +} diff --git a/src/js/api.js b/src/js/api.js index 579a1e1..4e37728 100644 --- a/src/js/api.js +++ b/src/js/api.js @@ -3,68 +3,82 @@ // ========================================== async function fetchModels(showInModal = true) { - const endpoint = showInModal - ? document.getElementById('apiEndpoint').value.trim() - : State.settings.apiEndpoint; - const apiKey = showInModal - ? document.getElementById('apiKey').value.trim() - : State.settings.apiKey; - - if (!endpoint || !apiKey) { - showToast('⚠️ Configure API endpoint and key first', 'warning'); - if (!showInModal) openSettings(); - return; - } - const btn = showInModal ? document.getElementById('fetchModelsBtn') : document.getElementById('quickFetchBtn'); - const originalText = btn.innerHTML; - btn.innerHTML = '⏳'; - btn.disabled = true; - + if (btn) { + btn.dataset.originalText = btn.innerHTML; + btn.innerHTML = '⏳'; + btn.disabled = true; + } + try { - const response = await fetch(endpoint.replace(/\/$/, '') + '/models', { - headers: { - 'Authorization': `Bearer ${apiKey}`, - 'Content-Type': 'application/json' + let models = []; + + if (Backend.isManaged) { + // Managed mode: fetch from backend aggregated models endpoint + const data = await Backend.listAllModels(); + models = (data.models || []).map(m => ({ + id: m.id || m.name, + owned_by: m.owned_by || m.provider || null, + config_id: m.config_id || null + })); + } else { + // Unmanaged mode: direct provider call + const endpoint = showInModal + ? document.getElementById('apiEndpoint').value.trim() + : State.settings.apiEndpoint; + const apiKey = showInModal + ? document.getElementById('apiKey').value.trim() + : State.settings.apiKey; + + if (!endpoint || !apiKey) { + showToast('⚠️ Configure API endpoint and key first', 'warning'); + if (!showInModal) openSettings(); + return; } - }); - - if (!response.ok) { - throw new Error(`API Error: ${response.status}`); + + const response = await fetch(endpoint.replace(/\/$/, '') + '/models', { + headers: { + 'Authorization': `Bearer ${apiKey}`, + 'Content-Type': 'application/json' + } + }); + + if (!response.ok) throw new Error(`API Error: ${response.status}`); + + const data = await response.json(); + const raw = data.data || data.models || data || []; + models = Array.isArray(raw) + ? raw.map(m => ({ + id: m.id || m.name || m, + owned_by: m.owned_by || null + })) + : []; } - - const data = await response.json(); - const models = data.data || data.models || data || []; - - State.models = Array.isArray(models) - ? models.map(m => ({ - id: m.id || m.name || m, - owned_by: m.owned_by || null - })).sort((a, b) => a.id.localeCompare(b.id)) - : []; - + + State.models = models.sort((a, b) => a.id.localeCompare(b.id)); saveModels(); - + if (showInModal) { populateModelSelect(document.getElementById('model')); } updateQuickModelSelector(); - showToast(`✅ Loaded ${State.models.length} models`, 'success'); } catch (error) { console.error('Fetch models error:', error); showToast(`❌ Failed: ${error.message}`, 'error'); } finally { - btn.innerHTML = originalText; - btn.disabled = false; + if (btn) { + btn.innerHTML = btn.dataset.originalText || '🔄'; + btn.disabled = false; + } } } async function sendApiRequest(messages) { const chatContainer = document.getElementById('chatMessages'); - + // Add typing indicator const typingDiv = document.createElement('div'); typingDiv.className = 'message assistant'; @@ -80,44 +94,60 @@ async function sendApiRequest(messages) { `; chatContainer.appendChild(typingDiv); scrollToBottom(); - + State.isGenerating = true; State.abortController = new AbortController(); document.getElementById('stopBtn').classList.add('visible'); document.getElementById('sendBtn').disabled = true; - - try { - const endpoint = State.settings.apiEndpoint.replace(/\/$/, '') + '/chat/completions'; - const model = document.getElementById('quickModel').value || State.settings.model; - - const response = await fetch(endpoint, { - method: 'POST', - headers: { - 'Content-Type': 'application/json', - 'Authorization': `Bearer ${State.settings.apiKey}` - }, - body: JSON.stringify({ - model, - messages: messages.map(m => ({ role: m.role, content: m.content })), - max_tokens: State.settings.maxTokens, - temperature: State.settings.temperature, - top_p: State.settings.topP, - presence_penalty: State.settings.presencePenalty, - stream: State.settings.stream - }), - signal: State.abortController.signal - }); - if (!response.ok) { - const error = await response.text(); - throw new Error(`API Error: ${response.status} - ${error}`); + const model = document.getElementById('quickModel').value || State.settings.model; + const userContent = messages[messages.length - 1]?.content || ''; + + try { + let response; + + if (Backend.isManaged) { + // ── Managed mode: route through backend proxy ── + // Backend loads conversation from DB, persists both messages. + response = await Backend.streamCompletion( + State.currentChatId, + userContent, + model, + State.abortController.signal + ); + } else { + // ── Unmanaged mode: direct provider call ── + const endpoint = State.settings.apiEndpoint.replace(/\/$/, '') + '/chat/completions'; + response = await fetch(endpoint, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + 'Authorization': `Bearer ${State.settings.apiKey}` + }, + body: JSON.stringify({ + model, + messages: messages.map(m => ({ role: m.role, content: m.content })), + max_tokens: State.settings.maxTokens, + temperature: State.settings.temperature, + top_p: State.settings.topP, + presence_penalty: State.settings.presencePenalty, + stream: State.settings.stream + }), + signal: State.abortController.signal + }); + + if (!response.ok) { + const error = await response.text(); + throw new Error(`API Error: ${response.status} - ${error}`); + } } document.getElementById('typingIndicator')?.remove(); - + let assistantContent = ''; - - if (State.settings.stream) { + const isStream = Backend.isManaged ? true : State.settings.stream; + + if (isStream) { const assistantDiv = document.createElement('div'); assistantDiv.className = 'message assistant'; assistantDiv.innerHTML = ` @@ -132,22 +162,22 @@ async function sendApiRequest(messages) { `; chatContainer.appendChild(assistantDiv); - + const reader = response.body.getReader(); const decoder = new TextDecoder(); - + while (true) { const { done, value } = await reader.read(); if (done) break; - + const chunk = decoder.decode(value); const lines = chunk.split('\n'); - + for (const line of lines) { if (line.startsWith('data: ')) { const data = line.slice(6); if (data === '[DONE]') continue; - + try { const parsed = JSON.parse(data); const content = parsed.choices?.[0]?.delta?.content || ''; @@ -162,26 +192,28 @@ async function sendApiRequest(messages) { const data = await response.json(); assistantContent = data.choices?.[0]?.message?.content || 'No response'; } - - messages.push({ - role: 'assistant', + + messages.push({ + role: 'assistant', content: assistantContent, timestamp: new Date().toISOString() }); - updateChat(State.currentChatId, { messages, model }); - - // Persist assistant response to backend (managed mode) - await onAssistantResponse(assistantContent, model); - - if (!State.settings.stream) { + + // In unmanaged mode, persist locally + to backend if connected + await updateChat(State.currentChatId, { messages, model }); + if (!Backend.isManaged) { + await onAssistantResponse(assistantContent, model); + } + + if (!isStream) { renderMessages(messages); } - + updateRegenerateButton(messages); - + } catch (error) { document.getElementById('typingIndicator')?.remove(); - + if (error.name === 'AbortError') { showToast('⚠️ Generation stopped', 'warning'); } else { diff --git a/src/js/app.js b/src/js/app.js index 89f20e6..b2f42d9 100644 --- a/src/js/app.js +++ b/src/js/app.js @@ -237,10 +237,13 @@ async function sendMessage() { if (!message || State.isGenerating) return; - if (!State.settings.apiEndpoint || !State.settings.apiKey) { - showToast('⚠️ Configure API settings first', 'warning'); - openSettings(); - return; + // In managed mode, backend has the keys. In unmanaged, user must configure. + if (!Backend.isManaged) { + if (!State.settings.apiEndpoint || !State.settings.apiKey) { + showToast('⚠️ Configure API settings first', 'warning'); + openSettings(); + return; + } } input.value = ''; @@ -270,8 +273,11 @@ async function sendMessage() { State.currentChatId = chat.id; } else { await updateChat(State.currentChatId, { messages }); - // Persist user message to backend - await persistMessage(State.currentChatId, 'user', message); + // In unmanaged mode, persist user message to backend (if connected) + // In managed mode, the completion proxy persists both messages + if (!Backend.isManaged) { + await persistMessage(State.currentChatId, 'user', message); + } } renderMessages(messages); diff --git a/src/js/backend.js b/src/js/backend.js index 82411c5..42a64fa 100644 --- a/src/js/backend.js +++ b/src/js/backend.js @@ -178,6 +178,89 @@ const Backend = { }); }, + // ── API Configs ───────────────────────── + + async listConfigs() { + return this._authedFetch('/api/v1/api-configs'); + }, + + async createConfig(name, provider, endpoint, apiKey, modelDefault) { + return this._authedFetch('/api/v1/api-configs', { + method: 'POST', + body: JSON.stringify({ + name, provider, endpoint, + api_key: apiKey, + model_default: modelDefault + }) + }); + }, + + async deleteConfig(configId) { + return this._authedFetch(`/api/v1/api-configs/${configId}`, { + method: 'DELETE' + }); + }, + + // ── Models ────────────────────────────── + + async listAllModels() { + return this._authedFetch('/api/v1/models'); + }, + + // ── Chat Completion (streaming) ───────── + // Returns the raw Response so the caller can read the SSE stream. + // The backend persists user + assistant messages automatically. + + async streamCompletion(chatId, content, model, signal) { + const url = this.baseUrl + '/api/v1/chat/completions'; + const body = { chat_id: chatId, content, stream: true }; + if (model) body.model = model; + + const resp = await fetch(url, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + 'Authorization': `Bearer ${this.accessToken}` + }, + body: JSON.stringify(body), + signal + }); + + if (resp.status === 401 && this.refreshToken) { + const refreshed = await this.refresh(); + if (refreshed) { + return fetch(url, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + 'Authorization': `Bearer ${this.accessToken}` + }, + body: JSON.stringify(body), + signal + }); + } + } + + if (!resp.ok) { + const errBody = await resp.json().catch(() => ({})); + throw new Error(errBody.error || `HTTP ${resp.status}`); + } + + return resp; + }, + + async syncCompletion(chatId, content, model) { + return this._authedFetch('/api/v1/chat/completions', { + method: 'POST', + body: JSON.stringify({ + chat_id: chatId, + content, + model, + stream: false + }) + }); + }, + // ── HTTP Layer ────────────────────────── async _authedFetch(path, opts = {}) {