Changeset 0.25.4 (#164)

This commit is contained in:
2026-03-09 20:17:46 +00:00
parent 2f7a0fb027
commit dbc1a97343
22 changed files with 660 additions and 491 deletions

View File

@@ -472,19 +472,21 @@ func (h *FileHandler) UploadToProject(c *gin.Context) {
userID := getUserID(c)
projectID := c.Param("id")
// Verify project access (user owns or is team member)
var exists bool
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT EXISTS(
SELECT 1 FROM projects
WHERE id = $1 AND (owner_id = $2 OR team_id IN (
SELECT team_id FROM team_members WHERE user_id = $2 AND is_active = true
))
)
`), projectID, userID).Scan(&exists)
if err != nil || !exists {
c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return
// Verify project access (user owns or is team member; admins bypass)
if c.GetString("role") != "admin" {
var exists bool
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT EXISTS(
SELECT 1 FROM projects
WHERE id = $1 AND (owner_id = $2 OR team_id IN (
SELECT team_id FROM team_members WHERE user_id = $2 AND is_active = true
))
)
`), projectID, userID).Scan(&exists)
if err != nil || !exists {
c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return
}
}
file, header, err := c.Request.FormFile("file")
@@ -551,19 +553,21 @@ func (h *FileHandler) ListByProject(c *gin.Context) {
userID := getUserID(c)
projectID := c.Param("id")
// Verify project access
var exists bool
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT EXISTS(
SELECT 1 FROM projects
WHERE id = $1 AND (owner_id = $2 OR team_id IN (
SELECT team_id FROM team_members WHERE user_id = $2 AND is_active = true
))
)
`), projectID, userID).Scan(&exists)
if err != nil || !exists {
c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return
// Admins bypass project ownership/membership checks
if c.GetString("role") != "admin" {
var exists bool
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT EXISTS(
SELECT 1 FROM projects
WHERE id = $1 AND (owner_id = $2 OR team_id IN (
SELECT team_id FROM team_members WHERE user_id = $2 AND is_active = true
))
)
`), projectID, userID).Scan(&exists)
if err != nil || !exists {
c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return
}
}
files, err := h.stores.Files.GetByProject(c.Request.Context(), projectID)

View File

@@ -4,7 +4,9 @@ import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"time"
@@ -37,12 +39,22 @@ type liveProviderConfig struct {
// defaultEndpoints maps provider names to their default API endpoints.
var defaultEndpoints = map[string]string{
"venice": "https://api.venice.ai/api/v1",
"openai": "https://api.openai.com/v1",
"anthropic": "https://api.anthropic.com/v1",
"venice": "https://api.venice.ai/api/v1",
"openai": "https://api.openai.com/v1",
"anthropic": "https://api.anthropic.com/v1",
"openrouter": "https://openrouter.ai/api/v1",
}
// requireLiveProvider resolves provider config from env vars and skips if not configured.
// providerKeyEnvs maps provider names to their API key env var names.
var providerKeyEnvs = map[string]string{
"venice": "VENICE_API_KEY",
"openai": "OPENAI_API_KEY",
"anthropic": "ANTHROPIC_API_KEY",
"openrouter": "OPENROUTER_API_KEY",
}
// requireLiveProvider resolves the primary provider config from env vars.
// Kept for backward compat — tests that only need one provider use this.
func requireLiveProvider(t *testing.T) liveProviderConfig {
t.Helper()
@@ -73,6 +85,191 @@ func requireLiveProvider(t *testing.T) liveProviderConfig {
return liveProviderConfig{Provider: provider, Key: key, Endpoint: endpoint}
}
// requireLiveProviders resolves all configured providers for failover tests.
// Reads LIVE_PROVIDERS (comma-separated, e.g. "venice,openai") and resolves
// API keys from {PROVIDER}_API_KEY env vars. Falls back to requireLiveProvider
// if LIVE_PROVIDERS is not set.
func requireLiveProviders(t *testing.T) []liveProviderConfig {
t.Helper()
list := os.Getenv("LIVE_PROVIDERS")
if list == "" {
// Fallback: just the primary provider
return []liveProviderConfig{requireLiveProvider(t)}
}
var configs []liveProviderConfig
for _, name := range splitTrim(list, ",") {
keyEnv := providerKeyEnvs[name]
if keyEnv == "" {
keyEnv = strings.ToUpper(name) + "_API_KEY"
}
key := os.Getenv(keyEnv)
if key == "" {
t.Logf(" Skipping provider %s: %s not set", name, keyEnv)
continue
}
endpoint := os.Getenv(strings.ToUpper(name) + "_API_URL")
if endpoint == "" {
endpoint = defaultEndpoints[name]
}
if endpoint == "" {
t.Logf(" Skipping provider %s: no endpoint", name)
continue
}
configs = append(configs, liveProviderConfig{Provider: name, Key: key, Endpoint: endpoint})
}
if len(configs) == 0 {
t.Skip("no live providers configured — set LIVE_PROVIDERS + API keys")
}
t.Logf(" Live providers: %d configured", len(configs))
return configs
}
// splitTrim splits s by sep and trims whitespace from each element.
func splitTrim(s, sep string) []string {
parts := strings.Split(s, sep)
var out []string
for _, p := range parts {
p = strings.TrimSpace(p)
if p != "" {
out = append(out, p)
}
}
return out
}
// providerModel holds a resolved (configID, modelID) pair from a provider.
type providerModel struct {
ConfigID string
ModelID string
Provider string
}
// setupAllProviders creates configs and enables a model for each provider.
// Tolerant: if a provider fails setup, it's skipped. Fails only if zero succeed.
func setupAllProviders(t *testing.T, h *testHarness, adminToken string, providerList []liveProviderConfig) []providerModel {
t.Helper()
var models []providerModel
for _, pc := range providerList {
configID, modelID, err := trySetupProvider(h, adminToken, pc)
if err != nil {
t.Logf(" ⚠ %s setup failed: %v — skipping", pc.Provider, err)
continue
}
t.Logf(" Provider %s ready, model %s enabled", pc.Provider, modelID)
models = append(models, providerModel{ConfigID: configID, ModelID: modelID, Provider: pc.Provider})
}
if len(models) == 0 {
t.Fatal("no providers available after setup — all failed")
}
return models
}
// trySetupProvider is like setupProviderWithModel but returns error instead of t.Fatal.
func trySetupProvider(h *testHarness, adminToken string, pc liveProviderConfig) (string, string, error) {
// Create provider
w := h.request("POST", "/api/v1/admin/configs", adminToken, map[string]interface{}{
"name": pc.Provider + " Test", "provider": pc.Provider,
"endpoint": pc.Endpoint, "api_key": pc.Key,
})
if w.Code != http.StatusCreated {
return "", "", fmt.Errorf("create config: %d: %s", w.Code, w.Body.String())
}
var cfg map[string]interface{}
decode(w, &cfg)
configID := cfg["id"].(string)
// Fetch models
w = h.request("POST", "/api/v1/admin/models/fetch", adminToken,
map[string]interface{}{"provider_config_id": configID})
if w.Code != http.StatusOK {
return "", "", fmt.Errorf("fetch models: %d: %s", w.Code, w.Body.String())
}
// Find and enable a non-reasoning model
w = h.request("GET", "/api/v1/admin/models", adminToken, nil)
var modelsResp map[string]interface{}
decode(w, &modelsResp)
var catalogID, modelID string
var fallbackCatalogID, fallbackModelID string
for _, raw := range modelsResp["models"].([]interface{}) {
m := raw.(map[string]interface{})
if m["visibility"].(string) != "disabled" {
continue
}
// Only pick models from this provider
if m["provider_config_id"] != nil && m["provider_config_id"].(string) != configID {
continue
}
mid := m["model_id"].(string)
cid := m["id"].(string)
if fallbackCatalogID == "" {
fallbackCatalogID = cid
fallbackModelID = mid
}
isReasoning := false
if caps, ok := m["capabilities"].(map[string]interface{}); ok {
if r, exists := caps["reasoning"]; exists && r == true {
isReasoning = true
}
}
if !isReasoning {
catalogID = cid
modelID = mid
break
}
}
if catalogID == "" {
catalogID = fallbackCatalogID
modelID = fallbackModelID
}
if catalogID == "" {
return "", "", fmt.Errorf("no disabled model found")
}
w = h.request("PUT", "/api/v1/admin/models/"+catalogID, adminToken,
map[string]interface{}{"visibility": "enabled"})
if w.Code != http.StatusOK {
return "", "", fmt.Errorf("enable model: %d: %s", w.Code, w.Body.String())
}
return configID, modelID, nil
}
// tryCompletion attempts a completion against each provider/model in order.
// Returns the first successful response and the provider that worked.
// Fails only if ALL providers fail.
func tryCompletion(t *testing.T, h *testHarness, token, channelID string, models []providerModel, stream bool) (*httptest.ResponseRecorder, providerModel) {
t.Helper()
var lastW *httptest.ResponseRecorder
for _, pm := range models {
w := h.request("POST", "/api/v1/chat/completions", token, map[string]interface{}{
"channel_id": channelID,
"content": "Say ok",
"model": pm.ModelID,
"provider_config_id": pm.ConfigID,
"stream": &stream,
"max_tokens": 1200,
})
if w.Code == http.StatusOK {
t.Logf(" ✓ Completion via %s/%s (status %d)", pm.Provider, pm.ModelID, w.Code)
return w, pm
}
t.Logf(" ✗ %s/%s returned %d — trying next", pm.Provider, pm.ModelID, w.Code)
lastW = w
}
// All failed
body := ""
if lastW != nil {
body = lastW.Body.String()
}
t.Fatalf("all %d providers failed; last: %s", len(models), body)
return nil, providerModel{} // unreachable
}
// setupProviderWithModel creates a provider config, fetches models, and enables
// the first available model. Returns (configID, enabledModelID).
func setupProviderWithModel(t *testing.T, h *testHarness, adminToken string, pc liveProviderConfig) (string, string) {
@@ -334,10 +531,10 @@ func TestLive_FetchModelsCapabilities(t *testing.T) {
// TestLive_ChatCompletion sends an actual non-streaming chat completion.
func TestLive_ChatCompletion(t *testing.T) {
h := setupHarness(t)
pc := requireLiveProvider(t)
liveProvs := requireLiveProviders(t)
_, adminToken := h.createAdminUser("admin", "admin@test.com")
configID, modelID := setupProviderWithModel(t, h, adminToken, pc)
models := setupAllProviders(t, h, adminToken, liveProvs)
w := h.request("POST", "/api/v1/channels", adminToken, map[string]interface{}{
"title": "Chat Test", "type": "direct",
@@ -349,18 +546,7 @@ func TestLive_ChatCompletion(t *testing.T) {
decode(w, &ch)
channelID := ch["id"].(string)
stream := false
w = h.request("POST", "/api/v1/chat/completions", adminToken, map[string]interface{}{
"channel_id": channelID,
"content": "Say ok",
"model": modelID,
"provider_config_id": configID,
"stream": &stream,
"max_tokens": 1200,
})
if w.Code != http.StatusOK {
t.Fatalf("completion: want 200, got %d: %s", w.Code, w.Body.String())
}
w, _ = tryCompletion(t, h, adminToken, channelID, models, false)
t.Logf(" ✓ Completion succeeded: %s", w.Body.String()[:min(200, w.Body.Len())])
}
@@ -368,10 +554,10 @@ func TestLive_ChatCompletion(t *testing.T) {
// creates a usage_log row with token counts.
func TestLive_UsageLogging(t *testing.T) {
h := setupHarness(t)
pc := requireLiveProvider(t)
liveProvs := requireLiveProviders(t)
_, adminToken := h.createAdminUser("admin", "admin@test.com")
configID, modelID := setupProviderWithModel(t, h, adminToken, pc)
models := setupAllProviders(t, h, adminToken, liveProvs)
w := h.request("POST", "/api/v1/channels", adminToken, map[string]interface{}{
"title": "Usage Test", "type": "direct",
@@ -382,23 +568,12 @@ func TestLive_UsageLogging(t *testing.T) {
var ch map[string]interface{}
decode(w, &ch)
stream := false
w = h.request("POST", "/api/v1/chat/completions", adminToken, map[string]interface{}{
"channel_id": ch["id"].(string),
"content": "Say ok",
"model": modelID,
"provider_config_id": configID,
"stream": &stream,
"max_tokens": 1200,
})
if w.Code != http.StatusOK {
t.Fatalf("completion: %d: %s", w.Code, w.Body.String())
}
_, used := tryCompletion(t, h, adminToken, ch["id"].(string), models, false)
var rowCount, promptTokens, completionTokens int
err := database.TestDB.QueryRow(
"SELECT COUNT(*), COALESCE(SUM(prompt_tokens), 0), COALESCE(SUM(completion_tokens), 0) FROM usage_log WHERE provider_config_id = "+database.PH(1),
configID).Scan(&rowCount, &promptTokens, &completionTokens)
used.ConfigID).Scan(&rowCount, &promptTokens, &completionTokens)
if err != nil {
t.Fatalf("query usage_log: %v", err)
}
@@ -408,16 +583,16 @@ func TestLive_UsageLogging(t *testing.T) {
if promptTokens == 0 {
t.Fatal("completion should report prompt tokens")
}
t.Logf(" ✓ Usage: %d row(s), prompt=%d completion=%d", rowCount, promptTokens, completionTokens)
t.Logf(" ✓ Usage: %d row(s), prompt=%d completion=%d (via %s)", rowCount, promptTokens, completionTokens, used.Provider)
}
// TestLive_StreamingUsageLogging verifies streaming completions log usage.
func TestLive_StreamingUsageLogging(t *testing.T) {
h := setupHarness(t)
pc := requireLiveProvider(t)
liveProvs := requireLiveProviders(t)
_, adminToken := h.createAdminUser("admin", "admin@test.com")
configID, modelID := setupProviderWithModel(t, h, adminToken, pc)
models := setupAllProviders(t, h, adminToken, liveProvs)
w := h.request("POST", "/api/v1/channels", adminToken, map[string]interface{}{
"title": "Stream Usage Test", "type": "direct",
@@ -428,23 +603,12 @@ func TestLive_StreamingUsageLogging(t *testing.T) {
var ch map[string]interface{}
decode(w, &ch)
stream := true
w = h.request("POST", "/api/v1/chat/completions", adminToken, map[string]interface{}{
"channel_id": ch["id"].(string),
"content": "Say ok",
"model": modelID,
"provider_config_id": configID,
"stream": &stream,
"max_tokens": 1200,
})
if w.Code != http.StatusOK {
t.Fatalf("streaming completion: %d: %s", w.Code, w.Body.String())
}
_, used := tryCompletion(t, h, adminToken, ch["id"].(string), models, true)
var rowCount, promptTokens, completionTokens int
err := database.TestDB.QueryRow(
"SELECT COUNT(*), COALESCE(SUM(prompt_tokens), 0), COALESCE(SUM(completion_tokens), 0) FROM usage_log WHERE provider_config_id = "+database.PH(1),
configID).Scan(&rowCount, &promptTokens, &completionTokens)
used.ConfigID).Scan(&rowCount, &promptTokens, &completionTokens)
if err != nil {
t.Fatalf("query usage_log: %v", err)
}
@@ -454,7 +618,7 @@ func TestLive_StreamingUsageLogging(t *testing.T) {
if promptTokens == 0 {
t.Fatal("streaming completion should report prompt tokens")
}
t.Logf(" ✓ Streaming usage: %d row(s), prompt=%d completion=%d", rowCount, promptTokens, completionTokens)
t.Logf(" ✓ Streaming usage: %d row(s), prompt=%d completion=%d (via %s)", rowCount, promptTokens, completionTokens, used.Provider)
}
// TestLive_PricingFromCatalog verifies model sync populates pricing.

View File

@@ -449,16 +449,20 @@ func (h *ProjectHandler) loadAndAuthorize(c *gin.Context) (*models.Project, bool
}
userID := getUserID(c)
teamIDs, _ := h.stores.Teams.GetUserTeamIDs(c.Request.Context(), userID)
ok, err := h.stores.Projects.UserCanAccess(c.Request.Context(), userID, projectID, teamIDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "access check failed"})
return nil, false
}
if !ok {
c.JSON(http.StatusNotFound, gin.H{"error": "project not found"})
return nil, false
// Admins can access all projects
if c.GetString("role") != "admin" {
teamIDs, _ := h.stores.Teams.GetUserTeamIDs(c.Request.Context(), userID)
ok, err := h.stores.Projects.UserCanAccess(c.Request.Context(), userID, projectID, teamIDs)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "access check failed"})
return nil, false
}
if !ok {
c.JSON(http.StatusNotFound, gin.H{"error": "project not found"})
return nil, false
}
}
project, err := h.stores.Projects.GetByID(c.Request.Context(), projectID)

View File

@@ -171,8 +171,11 @@ func (h *SettingsHandler) GetSettings(c *gin.Context) {
userID := getUserID(c)
var settingsRaw string
// Handle three states: SQL NULL, JSON null, or valid JSON object.
// NULLIF(settings, 'null') converts JSON null to SQL NULL,
// then COALESCE handles both SQL NULL cases.
err := database.DB.QueryRow(
database.Q(`SELECT settings::text FROM users WHERE id = $1`), userID,
database.Q(`SELECT COALESCE(NULLIF(settings::text, 'null'), '{}') FROM users WHERE id = $1`), userID,
).Scan(&settingsRaw)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load settings"})
@@ -202,12 +205,25 @@ func (h *SettingsHandler) UpdateSettings(c *gin.Context) {
return
}
// JSONB merge — existing keys preserved, incoming keys overwrite
// JSONB merge — existing keys preserved, incoming keys overwrite.
// Must handle three column states:
// 1. SQL NULL (never saved)
// 2. JSON literal null (corrupted by earlier bug)
// 3. JSON array (corrupted by repeated appends to null)
// 4. Valid JSON object (normal)
// NULLIF converts JSON null → SQL NULL, then COALESCE → '{}',
// and we enforce jsonb_typeof = 'object' to catch array corruption.
var mergeQuery string
if database.IsSQLite() {
mergeQuery = `UPDATE users SET settings = json_patch(settings, ?), updated_at = datetime('now') WHERE id = ?`
mergeQuery = `UPDATE users SET settings = json_patch(
CASE WHEN settings IS NULL OR settings = 'null' OR json_type(settings) != 'object'
THEN '{}' ELSE settings END,
?), updated_at = datetime('now') WHERE id = ?`
} else {
mergeQuery = `UPDATE users SET settings = settings || $1::jsonb, updated_at = NOW() WHERE id = $2`
mergeQuery = `UPDATE users SET settings = (
CASE WHEN settings IS NULL OR settings = 'null'::jsonb OR jsonb_typeof(settings) != 'object'
THEN '{}'::jsonb ELSE settings END
) || $1::jsonb, updated_at = NOW() WHERE id = $2`
}
_, err = database.DB.Exec(mergeQuery, string(patch), userID)
if err != nil {