Changeset 0.29.0 (#195)

This commit is contained in:
2026-03-17 16:28:47 +00:00
parent 128cbb8174
commit 5d637d3a90
129 changed files with 9418 additions and 3016 deletions

View File

@@ -1,5 +1,109 @@
# Changelog # Changelog
## [0.29.0] — 2026-03-17
### Summary
Starlark sandbox + permission model. Server-side extension runtime
with pre-completion filter chain, permission-gated modules, and
admin review workflow. Six changesets (CS0CS5) on top of Phase 0
store cleanup (12 changesets, CS0CS7b).
Phase 0 migrated ~242 raw SQL calls into the store interface layer,
giving the Starlark sandbox a clean API surface. Phase 1 builds the
runtime: sandboxed eval loop, permission model, secrets/notifications
modules, task executor integration, and pre-completion filter chain
with extension discovery.
**DB rebuild required.** Migration 016 rewritten in-place (adds
`status` column to `packages`, `extension_permissions` table).
### New
- **Pre-completion filter chain** (`server/filters/`) — composable
`PreCompletionFilter` interface with ordered execution, error
isolation, and logging. `Chain.Register()` / `Chain.Execute()`.
Built-in filters use order 099; extension filters use 100+.
- **KB auto-inject filter** — `KBInjectFilter` (order 10) replaces
inline `BuildKBHint()`. Scans persona, project, channel, and
personal KBs. Reference implementation for the filter model.
- **Starlark sandbox** (`server/sandbox/`) — `go.starlark.net`
interpreter with step limits (1M ops default), context timeout,
captured print output, disabled `load()`. `MakeModule()` helper
for exposing Go functions to Starlark scripts.
- **Sandbox runner** — `Runner.ExecPackage()` loads script from
manifest `_starlark_script`, assembles module set from granted
permissions. `CallEntryPoint()` for event-driven invocation.
- **Permission model** — `extension_permissions` table tracks
declared capabilities from package manifests. Admin grant/revoke
controls which modules the sandbox injects at runtime. Lifecycle:
`install → pending_review → grant-all → active`, revoke →
`suspended`. Auto-transitions on grant/revoke.
- **Extension permission constants** — `secrets.read`,
`notifications.send`, `filters.pre_completion`, `db.read`,
`db.write`, `api.http`. Validated against manifest declarations.
- **Secrets module** — `secrets.get(key)` / `secrets.list()`.
Per-package key-value store backed by GlobalConfig
(`ext_secrets:{packageID}`). Admin CRUD endpoints.
- **Notifications module** — `notifications.send(user_id, title,
body?, type?)`. Wraps platform notification service. Extensions
cannot send email or bypass user preferences.
- **Starlark filter** — `StarlarkFilter` bridges the filter chain
to extension scripts. Calls `on_pre_completion(ctx)`, parses
return `[{"role": "system", "content": "..."}]`.
- **Starlark filter discovery** — `DiscoverStarlarkFilters` scans
active packages with `filters.pre_completion` grant at startup.
- **`task_type: "starlark"`** — executor `executeStarlark` loads
package by `system_function` (package ID), calls `on_run()`
entry point. Output to channel/webhook/run record.
- **Extension secrets admin** — `GET/PUT/DELETE
/admin/extensions/:id/secrets`. GET returns keys only (not values).
- **Extension permission admin** — `GET .../permissions`,
`GET .../review`, `POST .../grant`, `POST .../revoke`,
`POST .../grant-all`.
- **ICD runner packaging tier** — 18 tests: permission lifecycle
(install → pending_review → grant → active → revoke → suspended),
secrets CRUD, invalid permission rejection.
### Changed
- **`BuildKBHint` deprecated** — replaced by `KBInjectFilter` in
the pre-completion filter chain. Kept for backward compatibility;
removal scheduled for v0.30.0.
- **`packages` table** — `status` column added (`active`,
`pending_review`, `suspended`). `PackageStore.SetStatus()` method.
- **`CompletionHandler`** — `filterChain` field + `SetFilterChain()`
setter. Filter chain runs between system prompts and message history.
- **`Executor`** — `runner` field + `SetRunner()`. Starlark dispatch
branch between action and prompt types.
- **`AdminInstallExtension`** — calls `SyncManifestPermissions()`
after package creation. Parses `"permissions"` from manifest.
- **Starlark task validation** — `task.starlark` RBAC gate enforced.
`system_function` required (holds package ID). Package must exist
and be starlark tier.
### Phase 0 — Store Cleanup
- **~242 raw SQL calls eliminated** — every handler, tool, and
middleware file now goes through the store interface.
- **Constructor signature changes** —
`enforcePrivateProviderPolicy(ctx, stores, ...)`,
`ResolveProviderConfig(stores, vault, ...)`,
`NewChannelHandler(stores)`.
- **30+ new store methods** across NoteStore, MessageStore, FileStore,
ProjectStore, TeamStore, ProviderStore, CatalogStore,
GlobalConfigStore, AuditStore, ChannelStore.
### Fixed
- **`API._delete` alias** — `delete` is a JS reserved word; added
`_delete` as alias for `_del` on the API client.
- **`safeString` nil handling** — returns Go `nil`, not `""`.
- **nil JSONMap → PG jsonb** — `ToolCalls`/`Metadata` initialized
to `models.JSONMap{}` before store calls.
- **`GetByID` vs `GetParentAndRole`** — cursor update uses
`GetParentAndRole` (matches original raw SQL deleted_at filter).
## [0.28.8] — 2026-03-15 ## [0.28.8] — 2026-03-15
### Summary ### Summary

View File

@@ -1 +1 @@
0.28.8 0.29.0

View File

@@ -30,7 +30,7 @@ v0.9.xv0.28.7 Foundation through Platform Polish ✅
│ │ │ │
Extension Track Operations Track Extension Track Operations Track
│ │ │ │
v0.29.0 Starlark Sandbox v0.32.0 Multi-Replica HA v0.29.0 Starlark Sandbox v0.32.0 Multi-Replica HA
v0.29.1 API Extensions v0.33.0 Observability v0.29.1 API Extensions v0.33.0 Observability
v0.29.2 DB Extensions v0.34.0 Data Portability v0.29.2 DB Extensions v0.34.0 Data Portability
v0.29.3 Workflow Forms │ v0.29.3 Workflow Forms │
@@ -50,12 +50,6 @@ v0.9.xv0.28.7 Foundation through Platform Polish ✅
STT/TTS, desktop app) STT/TTS, desktop app)
``` ```
**Target deployment (v0.50.0):** 3-node cluster, 3-node PG (patroni/cnpg),
S3 (Ceph RGW), CephFS. 510 teams × ~5 users, 4 admins, 100+ anonymous
visitors. Multi-replica backend for node-level HA.
---
## Completed: v0.28.0 — Platform Polish ## Completed: v0.28.0 — Platform Polish
Audit arc, frontend decomposition, security, infrastructure. Eight Audit arc, frontend decomposition, security, infrastructure. Eight
@@ -132,32 +126,50 @@ Memory profile (measured on cluster via `kubectl top`):
Sequential. Each version builds on the previous. Delivers the package Sequential. Each version builds on the previous. Delivers the package
ecosystem, workflow capabilities, and SDK-based surface architecture. ecosystem, workflow capabilities, and SDK-based surface architecture.
### v0.29.0 — Starlark Sandbox + Permission Model ### v0.29.0 — Starlark Sandbox + Permission Model
Server-side extension runtime. Prove the eval loop, permission pipeline, Server-side extension runtime. Eval loop, permission pipeline,
and admin UI before adding capabilities. pre-completion filter chain, and admin review workflow.
Depends on: v0.28.8. Depends on: v0.28.8.
**Phase 0 — Store cleanup (prerequisite):** **Phase 0 — Store cleanup (prerequisite):**
- [ ] `SELECT *` → explicit column lists in both store packages - [x] Raw SQL hunt: all ~242 `database.DB.*` calls outside `store/`
- [ ] Raw SQL hunt: all `database.DB.*` outside `store/` → store methods migrated to store interface methods (CS0CS7b, 12 changesets)
- [ ] CI gate: grep-enforced, zero hits - [x] CI green on both PG and SQLite pipelines
- [x] ICD runner: 579/580 pass, 1 expected skip
- [x] Documented exception: `events/pg_broadcast.go` (`pg_notify`,
PG-only, no store abstraction needed)
**Starlark runtime:** **Phase 1 — Starlark runtime:**
- [ ] `go.starlark.net` integration (eval loop, timeout, memory ceiling) - [x] Pre-completion filter chain: composable `PreCompletionFilter`
- [ ] Permission model: manifest declarations, admin grant/revoke, interface + `Chain` registry. KB auto-inject refactored as first
`extension_permissions` table built-in filter. Extension filters register at order 100+ (CS0)
- [ ] Runtime enforcement: sandbox injects only granted modules - [x] `go.starlark.net` integration: sandboxed eval with step limits
- [ ] Admin UI: permission review, grant/revoke, audit log. (1M ops default), context timeout, captured print output,
Task permission management (deferred from v0.28.7). disabled `load()`. `MakeModule` helper for Go→Starlark (CS1)
- [ ] Extension lifecycle: `install → pending_review → approved → active` - [x] Permission model: `extension_permissions` table (in 016),
- [ ] Initial modules: `secrets` (vault-backed), `notifications` `status` column on `packages` (`active`/`pending_review`/
- [ ] `task_type: "starlark"`: sandbox execution with task context. `suspended`). Manifest `"permissions"` array parsed on install.
RBAC gate from v0.28.7 enforced. Admin review, grant, revoke, grant-all endpoints (CS2)
- [ ] KB auto-injection: server-side pre-completion filter (Go built-in). - [x] Runtime enforcement: `Runner.buildModules()` injects only
Reference implementation for the filter model Starlark mirrors. granted modules into sandbox namespace (CS3)
- [ ] ICD runner: `packaging` test tier (deferred from v0.28.7) - [x] Extension lifecycle: `install → pending_review → grant-all →
active`, revoke → `suspended`. Auto-transitions on grant/revoke.
- [x] Initial modules: `secrets` (GlobalConfig-backed, per-package
key-value store, admin CRUD), `notifications` (wraps
notification service, `send(user_id, title, body?, type?)`) (CS3)
- [x] `task_type: "starlark"`: executor `executeStarlark` loads
package by ID, calls `on_run()` entry point via runner.
RBAC gate `task.starlark` enforced. `system_function` field
holds package ID (CS4)
- [x] KB auto-injection: server-side pre-completion filter chain.
Reference implementation for the filter model Starlark
extensions mirror via `on_pre_completion(ctx)` (CS0+CS3)
- [x] Starlark filter discovery: `DiscoverStarlarkFilters` scans
active packages with `filters.pre_completion` grant (CS3)
- [x] ICD runner: `packaging` test tier — 18 tests covering
permission lifecycle + secrets CRUD (CS5)
### v0.29.1 — API Extensions ### v0.29.1 — API Extensions

View File

@@ -14,7 +14,9 @@
* 7. tier-authz.js — Permission boundary tests * 7. tier-authz.js — Permission boundary tests
* 8. tier-security.js — Adversarial red-team tests (auth, cross-tenant, input validation) * 8. tier-security.js — Adversarial red-team tests (auth, cross-tenant, input validation)
* 9. tier-providers.js — Three-tier provider CRUD + live completions * 9. tier-providers.js — Three-tier provider CRUD + live completions
* 10. ui.js — Render functions, export, provider setup panel * 10. tier-packaging.js — Extension permission lifecycle + secrets (v0.29.0)
* 11. tier-sdk.js — Switchboard SDK contract tests
* 12. ui.js — Render functions, export, provider setup panel
* 11. (this file) — Boot * 11. (this file) — Boot
*/ */
(function () { (function () {
@@ -81,6 +83,7 @@
'tier-authz.js', 'tier-authz.js',
'tier-security.js', 'tier-security.js',
'tier-providers.js', 'tier-providers.js',
'tier-packaging.js',
'tier-sdk.js', 'tier-sdk.js',
'ui.js' 'ui.js'
]; ];

View File

@@ -0,0 +1,182 @@
/**
* ICD Test Runner — Packaging Tier
*
* v0.29.0 CS5: Tests the extension permission lifecycle
* (install with permissions → pending_review → grant → active →
* revoke → suspended) and extension secrets CRUD.
*
* Requires admin token (from fixtures).
*/
(function () {
'use strict';
var T = window.ICD;
if (!T) return;
T.runPackaging = async function () {
var extId = 'icd-pkg-' + Date.now();
var extInstalled = false;
// ── Install with permissions → pending_review ──
await T.test('crud', 'packaging', 'POST /admin/extensions (with permissions)', async function () {
var d = await T.apiPost('/admin/extensions', {
ext_id: extId,
name: 'ICD Packaging Test',
version: '1.0.0',
tier: 'starlark',
description: 'Tests permission lifecycle',
author: 'icd-runner',
manifest: {
permissions: ['secrets.read', 'notifications.send'],
_starlark_script: 'def on_run():\n return "ok"'
},
is_enabled: true,
is_system: false
});
T.assertHasKey(d, 'data', 'install response');
T.assert(d.data.id === extId, 'id should match ext_id');
T.assert(d.data.tier === 'starlark', 'tier should be starlark');
extInstalled = true;
T.registerCleanup(function () {
if (extInstalled) return T.safeDelete('/admin/extensions/' + extId);
});
});
if (!extInstalled) return;
// ── Verify pending_review status ──
await T.test('crud', 'packaging', 'GET /admin/extensions/:id (status = pending_review)', async function () {
var d = await T.apiGet('/admin/extensions');
T.assertHasKey(d, 'data', 'admin list');
var pkg = d.data.find(function (e) { return e.id === extId; });
T.assert(pkg, 'installed package should be in admin list');
T.assert(pkg.status === 'pending_review', 'status should be pending_review, got: ' + pkg.status);
});
// ── List declared permissions ──
await T.test('crud', 'packaging', 'GET /admin/extensions/:id/permissions (declared)', async function () {
var d = await T.apiGet('/admin/extensions/' + extId + '/permissions');
T.assertHasKey(d, 'data', 'permissions response');
T.assert(Array.isArray(d.data), 'permissions should be array');
T.assert(d.data.length === 2, 'should have 2 declared permissions, got: ' + d.data.length);
var perms = d.data.map(function (p) { return p.permission; }).sort();
T.assert(perms[0] === 'notifications.send', 'first perm should be notifications.send');
T.assert(perms[1] === 'secrets.read', 'second perm should be secrets.read');
var allUngrated = d.data.every(function (p) { return p.granted === false; });
T.assert(allUngrated, 'all permissions should be ungranted initially');
});
// ── Review package ──
await T.test('crud', 'packaging', 'GET /admin/extensions/:id/review', async function () {
var d = await T.apiGet('/admin/extensions/' + extId + '/review');
T.assertHasKey(d, 'package', 'review response package');
T.assertHasKey(d, 'permissions', 'review response permissions');
T.assert(d.package.id === extId, 'review package id should match');
T.assert(Array.isArray(d.permissions), 'review permissions should be array');
T.assert(d.permissions.length === 2, 'review should show 2 permissions');
});
// ── Grant single permission ──
await T.test('crud', 'packaging', 'POST .../permissions/secrets.read/grant', async function () {
var d = await T.apiPost('/admin/extensions/' + extId + '/permissions/secrets.read/grant', {});
T.assert(d.status === 'granted', 'should return granted status');
T.assert(d.permission === 'secrets.read', 'should echo permission');
});
// ── Still pending_review (not all granted) ──
await T.test('crud', 'packaging', 'status still pending_review (1 of 2 granted)', async function () {
var d = await T.apiGet('/admin/extensions');
var pkg = d.data.find(function (e) { return e.id === extId; });
T.assert(pkg.status === 'pending_review', 'should still be pending_review, got: ' + pkg.status);
});
// ── Grant all remaining ──
await T.test('crud', 'packaging', 'POST .../permissions/grant-all', async function () {
var d = await T.apiPost('/admin/extensions/' + extId + '/permissions/grant-all', {});
T.assert(d.status === 'granted_all', 'should return granted_all status');
});
// ── Verify active status ──
await T.test('crud', 'packaging', 'status → active (all permissions granted)', async function () {
var d = await T.apiGet('/admin/extensions');
var pkg = d.data.find(function (e) { return e.id === extId; });
T.assert(pkg.status === 'active', 'should be active after grant-all, got: ' + pkg.status);
});
// ── Verify all granted ──
await T.test('crud', 'packaging', 'GET .../permissions (all granted)', async function () {
var d = await T.apiGet('/admin/extensions/' + extId + '/permissions');
var allGranted = d.data.every(function (p) { return p.granted === true; });
T.assert(allGranted, 'all permissions should be granted');
var hasGrantedBy = d.data.every(function (p) { return p.granted_by !== null; });
T.assert(hasGrantedBy, 'all permissions should have granted_by');
});
// ── Secrets CRUD ──
await T.test('crud', 'packaging', 'PUT /admin/extensions/:id/secrets (set)', async function () {
var d = await T.apiPut('/admin/extensions/' + extId + '/secrets', {
secrets: { api_key: 'sk-test-123', webhook_token: 'tok-abc' }
});
T.assert(d.status === 'saved', 'should return saved status');
T.assert(d.key_count === 2, 'should report 2 keys');
});
await T.test('crud', 'packaging', 'GET /admin/extensions/:id/secrets (keys only)', async function () {
var d = await T.apiGet('/admin/extensions/' + extId + '/secrets');
T.assertHasKey(d, 'data', 'secrets response');
T.assert(d.data.package_id === extId, 'package_id should match');
T.assert(Array.isArray(d.data.keys), 'keys should be array');
T.assert(d.data.keys.length === 2, 'should have 2 keys');
// Values should NOT be returned
T.assert(!d.data.api_key, 'raw values should not be exposed');
});
await T.test('crud', 'packaging', 'DELETE /admin/extensions/:id/secrets', async function () {
var d = await T.apiDelete('/admin/extensions/' + extId + '/secrets');
T.assert(d.status === 'deleted', 'should return deleted status');
});
await T.test('crud', 'packaging', 'GET .../secrets (empty after delete)', async function () {
var d = await T.apiGet('/admin/extensions/' + extId + '/secrets');
T.assert(d.data.keys.length === 0, 'keys should be empty after delete');
});
// ── Revoke → suspended ──
await T.test('crud', 'packaging', 'POST .../permissions/secrets.read/revoke', async function () {
var d = await T.apiPost('/admin/extensions/' + extId + '/permissions/secrets.read/revoke', {});
T.assert(d.status === 'revoked', 'should return revoked status');
});
await T.test('crud', 'packaging', 'status → suspended (permission revoked)', async function () {
var d = await T.apiGet('/admin/extensions');
var pkg = d.data.find(function (e) { return e.id === extId; });
T.assert(pkg.status === 'suspended', 'should be suspended after revoke, got: ' + pkg.status);
});
// ── Invalid permission ──
await T.test('crud', 'packaging', 'POST .../permissions/bogus/grant (400)', async function () {
var d = await T.authFetch(API.accessToken, 'POST', '/admin/extensions/' + extId + '/permissions/bogus/grant', {});
T.assertStatus(d, 400, 'invalid permission should 400');
});
// ── Cleanup (DELETE handled by registerCleanup) ──
await T.test('crud', 'packaging', 'DELETE /admin/extensions/:id (cleanup)', async function () {
var d = await T.apiDelete('/admin/extensions/' + extId);
extInstalled = false;
});
};
})();

View File

@@ -238,6 +238,13 @@
}, 'SDK'); }, 'SDK');
T.el.controls.appendChild(btnSdk); T.el.controls.appendChild(btnSdk);
// Packaging tier button — extension permission lifecycle (v0.29.0)
var btnPkg = $('button', {
className: 'btn-secondary',
onClick: function () { T.runSuite('packaging'); }
}, 'Packaging');
T.el.controls.appendChild(btnPkg);
T.el.controls.appendChild(btnAll); T.el.controls.appendChild(btnAll);
T.el.controls.appendChild(btnClear); T.el.controls.appendChild(btnClear);
@@ -546,6 +553,7 @@
if (which === 'security' || (which === 'all' && T.fixtures.ready)) await T.runSecurity(); if (which === 'security' || (which === 'all' && T.fixtures.ready)) await T.runSecurity();
if (which === 'provider' || (which === 'all' && T.providerSetup.configured)) await T.runProviders(); if (which === 'provider' || (which === 'all' && T.providerSetup.configured)) await T.runProviders();
if (which === 'sdk' || which === 'all') { if (typeof T.runSdk === 'function') await T.runSdk(); } if (which === 'sdk' || which === 'all') { if (typeof T.runSdk === 'function') await T.runSdk(); }
if (which === 'packaging' || which === 'all') { if (typeof T.runPackaging === 'function') await T.runPackaging(); }
} catch (e) { } catch (e) {
T.results.push({ tier: '?', domain: 'runner', name: 'FATAL', status: 'fail', duration: 0, detail: String(e) }); T.results.push({ tier: '?', domain: 'runner', name: 'FATAL', status: 'fail', duration: 0, detail: String(e) });
} }

View File

@@ -8,7 +8,6 @@ import (
"strings" "strings"
"time" "time"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/providers" "git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/roles" "git.gobha.me/xcaliber/chat-switchboard/roles"
@@ -173,16 +172,22 @@ Write the summary in a way that would allow the conversation to continue natural
metaJSON, _ := json.Marshal(metadata) metaJSON, _ := json.Marshal(metadata)
siblingIdx := treepath.NextSiblingIndex(req.ChannelID, &lastMessageID) siblingIdx := treepath.NextSiblingIndex(req.ChannelID, &lastMessageID)
var summaryMsgID string summaryMsg := &models.Message{
err = database.DB.QueryRow(` ChannelID: req.ChannelID,
INSERT INTO messages (channel_id, parent_id, role, content, model, metadata, sibling_index, participant_type) ParentID: &lastMessageID,
VALUES ($1, $2, 'assistant', $3, $4, $5, $6, 'system') Role: "assistant",
RETURNING id Content: result.Content,
`, req.ChannelID, lastMessageID, result.Content, result.Model, string(metaJSON), siblingIdx).Scan(&summaryMsgID) Model: result.Model,
if err != nil { Metadata: models.JSONMap{},
SiblingIndex: siblingIdx,
ParticipantType: "system",
}
_ = json.Unmarshal(metaJSON, &summaryMsg.Metadata)
if err := s.stores.Messages.Create(ctx, summaryMsg); err != nil {
log.Printf("⚠ Failed to persist summary for channel %s: %v", req.ChannelID, err) log.Printf("⚠ Failed to persist summary for channel %s: %v", req.ChannelID, err)
return nil, fmt.Errorf("failed to save summary: %w", err) return nil, fmt.Errorf("failed to save summary: %w", err)
} }
summaryMsgID := summaryMsg.ID
// Update cursor to point to the summary node // Update cursor to point to the summary node
if err := treepath.UpdateCursor(req.ChannelID, req.UserID, summaryMsgID); err != nil { if err := treepath.UpdateCursor(req.ChannelID, req.UserID, summaryMsgID); err != nil {
@@ -268,11 +273,8 @@ func (s *Service) CheckRateLimit(ctx context.Context, userID string) error {
// GetUserTeamID returns the user's first team ID (for role resolution). // GetUserTeamID returns the user's first team ID (for role resolution).
func (s *Service) GetUserTeamID(ctx context.Context, userID string) *string { func (s *Service) GetUserTeamID(ctx context.Context, userID string) *string {
var teamID string teamID, _ := s.stores.Teams.GetFirstTeamIDForUser(ctx, userID)
err := database.DB.QueryRow(` if teamID == "" {
SELECT team_id FROM team_members WHERE user_id = $1 LIMIT 1
`, userID).Scan(&teamID)
if err != nil || teamID == "" {
return nil return nil
} }
return &teamID return &teamID

View File

@@ -2,12 +2,10 @@ package compaction
import ( import (
"context" "context"
"encoding/json"
"log" "log"
"sync" "sync"
"time" "time"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/roles" "git.gobha.me/xcaliber/chat-switchboard/roles"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
@@ -236,47 +234,12 @@ func (sc *Scanner) findCandidates(ctx context.Context) []models.Channel {
activityGap := time.Now().Add(-candidateActivityGap) activityGap := time.Now().Add(-candidateActivityGap)
maxAge := time.Now().Add(-candidateMaxAge) maxAge := time.Now().Add(-candidateMaxAge)
rows, err := database.DB.QueryContext(ctx, ` candidates, err := sc.stores.Channels.FindCompactionCandidates(ctx,
SELECT c.id, c.user_id, COALESCE(c.model, ''), COALESCE(c.settings::text, '{}'), activityGap, maxAge, candidateMinMessages, candidateMinChars, candidateBatchSize)
COUNT(m.id) AS msg_count,
COALESCE(SUM(LENGTH(m.content)), 0) AS total_chars
FROM channels c
JOIN messages m ON m.channel_id = c.id AND m.deleted_at IS NULL
WHERE c.type = 'direct'
AND c.is_archived = false
AND c.updated_at < $1
AND c.updated_at > $2
GROUP BY c.id
HAVING COUNT(m.id) >= $3
AND COALESCE(SUM(LENGTH(m.content)), 0) > $4
ORDER BY c.updated_at DESC
LIMIT $5
`, activityGap,
maxAge,
candidateMinMessages,
candidateMinChars,
candidateBatchSize,
)
if err != nil { if err != nil {
log.Printf("⚠ compaction: candidate query failed: %v", err) log.Printf("⚠ compaction: candidate query failed: %v", err)
return nil return nil
} }
defer rows.Close()
var candidates []models.Channel
for rows.Next() {
var ch models.Channel
var settingsRaw string
var msgCount, totalChars int
if err := rows.Scan(&ch.ID, &ch.UserID, &ch.Model, &settingsRaw, &msgCount, &totalChars); err != nil {
log.Printf("⚠ compaction: candidate scan: %v", err)
continue
}
if settingsRaw != "" {
_ = json.Unmarshal([]byte(settingsRaw), &ch.Settings)
}
candidates = append(candidates, ch)
}
return candidates return candidates
} }

View File

@@ -6,6 +6,7 @@
-- extension, or both. -- extension, or both.
-- --
-- v0.28.7: packages table + package_user_settings in 012. -- v0.28.7: packages table + package_user_settings in 012.
-- v0.29.0: status column + extension_permissions table.
-- ========================================== -- ==========================================
CREATE TABLE IF NOT EXISTS packages ( CREATE TABLE IF NOT EXISTS packages (
@@ -25,6 +26,8 @@ CREATE TABLE IF NOT EXISTS packages (
installed_by UUID REFERENCES users(id) ON DELETE SET NULL, installed_by UUID REFERENCES users(id) ON DELETE SET NULL,
manifest JSONB NOT NULL DEFAULT '{}', manifest JSONB NOT NULL DEFAULT '{}',
enabled BOOLEAN NOT NULL DEFAULT true, enabled BOOLEAN NOT NULL DEFAULT true,
status TEXT NOT NULL DEFAULT 'active'
CHECK (status IN ('active', 'pending_review', 'suspended')),
source TEXT NOT NULL DEFAULT 'core' source TEXT NOT NULL DEFAULT 'core'
CHECK (source IN ('core', 'builtin', 'extension')), CHECK (source IN ('core', 'builtin', 'extension')),
installed_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), installed_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
@@ -35,12 +38,14 @@ CREATE INDEX IF NOT EXISTS idx_packages_type ON packages(type);
CREATE INDEX IF NOT EXISTS idx_packages_enabled ON packages(enabled) WHERE enabled = true; CREATE INDEX IF NOT EXISTS idx_packages_enabled ON packages(enabled) WHERE enabled = true;
CREATE INDEX IF NOT EXISTS idx_packages_team ON packages(team_id); CREATE INDEX IF NOT EXISTS idx_packages_team ON packages(team_id);
CREATE INDEX IF NOT EXISTS idx_packages_source ON packages(source); CREATE INDEX IF NOT EXISTS idx_packages_source ON packages(source);
CREATE INDEX IF NOT EXISTS idx_packages_status ON packages(status);
COMMENT ON TABLE packages IS 'Unified package registry. Surfaces, extensions, and full packages. Replaces surface_registry + extensions tables.'; COMMENT ON TABLE packages IS 'Unified package registry. Surfaces, extensions, and full packages. Replaces surface_registry + extensions tables.';
COMMENT ON COLUMN packages.id IS 'Slug identifier from manifest "id" field. Used in URLs: /s/:id'; COMMENT ON COLUMN packages.id IS 'Slug identifier from manifest "id" field. Used in URLs: /s/:id';
COMMENT ON COLUMN packages.type IS 'surface = routable page, extension = hooks/tools/pipes, full = both'; COMMENT ON COLUMN packages.type IS 'surface = routable page, extension = hooks/tools/pipes, full = both';
COMMENT ON COLUMN packages.source IS 'core = page-engine seeded, builtin = extensions/builtin/ seeded, extension = admin-uploaded .pkg'; COMMENT ON COLUMN packages.source IS 'core = page-engine seeded, builtin = extensions/builtin/ seeded, extension = admin-uploaded .pkg';
COMMENT ON COLUMN packages.enabled IS 'Admin toggle — disabled surfaces redirect to / and hide from nav'; COMMENT ON COLUMN packages.enabled IS 'Admin toggle — disabled surfaces redirect to / and hide from nav';
COMMENT ON COLUMN packages.status IS 'Lifecycle: active (running), pending_review (needs admin permission grant), suspended (permission revoked)';
-- ========================================= -- =========================================
@@ -57,3 +62,29 @@ CREATE TABLE IF NOT EXISTS package_user_settings (
is_enabled BOOLEAN NOT NULL DEFAULT true, is_enabled BOOLEAN NOT NULL DEFAULT true,
PRIMARY KEY (package_id, user_id) PRIMARY KEY (package_id, user_id)
); );
-- =========================================
-- EXTENSION PERMISSIONS
-- =========================================
-- Declared capabilities from package manifests. Admin grants control
-- which modules the Starlark sandbox injects at runtime.
-- v0.29.0: Permission model for extension sandboxing.
CREATE TABLE IF NOT EXISTS extension_permissions (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
package_id TEXT NOT NULL REFERENCES packages(id) ON DELETE CASCADE,
permission TEXT NOT NULL,
granted BOOLEAN NOT NULL DEFAULT false,
granted_by UUID REFERENCES users(id) ON DELETE SET NULL,
granted_at TIMESTAMPTZ,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
UNIQUE(package_id, permission)
);
CREATE INDEX IF NOT EXISTS idx_ext_perm_package ON extension_permissions(package_id);
CREATE INDEX IF NOT EXISTS idx_ext_perm_granted ON extension_permissions(granted) WHERE granted = true;
COMMENT ON TABLE extension_permissions IS 'Declared permissions from package manifests. Admin grants control runtime module injection.';
COMMENT ON COLUMN extension_permissions.permission IS 'Capability key: secrets.read, notifications.send, filters.pre_completion, db.read, db.write, api.http';
COMMENT ON COLUMN extension_permissions.granted IS 'Admin has reviewed and approved this capability for the package';

View File

@@ -1,6 +1,7 @@
-- Chat Switchboard — 016 Packages (SQLite) -- Chat Switchboard — 016 Packages (SQLite)
-- Unified package registry. Replaces surface_registry + extensions. -- Unified package registry. Replaces surface_registry + extensions.
-- v0.28.7 -- v0.28.7: packages + package_user_settings
-- v0.29.0: status column + extension_permissions table
CREATE TABLE IF NOT EXISTS packages ( CREATE TABLE IF NOT EXISTS packages (
id TEXT PRIMARY KEY, id TEXT PRIMARY KEY,
@@ -19,6 +20,8 @@ CREATE TABLE IF NOT EXISTS packages (
installed_by TEXT REFERENCES users(id) ON DELETE SET NULL, installed_by TEXT REFERENCES users(id) ON DELETE SET NULL,
manifest TEXT NOT NULL DEFAULT '{}', manifest TEXT NOT NULL DEFAULT '{}',
enabled INTEGER NOT NULL DEFAULT 1, enabled INTEGER NOT NULL DEFAULT 1,
status TEXT NOT NULL DEFAULT 'active'
CHECK (status IN ('active', 'pending_review', 'suspended')),
source TEXT NOT NULL DEFAULT 'core' source TEXT NOT NULL DEFAULT 'core'
CHECK (source IN ('core', 'builtin', 'extension')), CHECK (source IN ('core', 'builtin', 'extension')),
installed_at TEXT NOT NULL DEFAULT (datetime('now')), installed_at TEXT NOT NULL DEFAULT (datetime('now')),
@@ -29,10 +32,10 @@ CREATE INDEX IF NOT EXISTS idx_packages_type ON packages(type);
CREATE INDEX IF NOT EXISTS idx_packages_enabled ON packages(enabled); CREATE INDEX IF NOT EXISTS idx_packages_enabled ON packages(enabled);
CREATE INDEX IF NOT EXISTS idx_packages_team ON packages(team_id); CREATE INDEX IF NOT EXISTS idx_packages_team ON packages(team_id);
CREATE INDEX IF NOT EXISTS idx_packages_source ON packages(source); CREATE INDEX IF NOT EXISTS idx_packages_source ON packages(source);
CREATE INDEX IF NOT EXISTS idx_packages_status ON packages(status);
-- Per-user overrides for packages (enable/disable, custom config). -- Per-user overrides for packages (enable/disable, custom config).
-- Lives in 016 (not 012) because of FK dependency on packages table.
CREATE TABLE IF NOT EXISTS package_user_settings ( CREATE TABLE IF NOT EXISTS package_user_settings (
package_id TEXT NOT NULL REFERENCES packages(id) ON DELETE CASCADE, package_id TEXT NOT NULL REFERENCES packages(id) ON DELETE CASCADE,
@@ -41,3 +44,20 @@ CREATE TABLE IF NOT EXISTS package_user_settings (
is_enabled INTEGER NOT NULL DEFAULT 1, is_enabled INTEGER NOT NULL DEFAULT 1,
PRIMARY KEY (package_id, user_id) PRIMARY KEY (package_id, user_id)
); );
-- Extension permissions (v0.29.0)
CREATE TABLE IF NOT EXISTS extension_permissions (
id TEXT PRIMARY KEY,
package_id TEXT NOT NULL REFERENCES packages(id) ON DELETE CASCADE,
permission TEXT NOT NULL,
granted INTEGER NOT NULL DEFAULT 0,
granted_by TEXT REFERENCES users(id) ON DELETE SET NULL,
granted_at TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
UNIQUE(package_id, permission)
);
CREATE INDEX IF NOT EXISTS idx_ext_perm_package ON extension_permissions(package_id);
CREATE INDEX IF NOT EXISTS idx_ext_perm_granted ON extension_permissions(granted);

View File

@@ -0,0 +1,83 @@
// Package filters — discover.go
//
// v0.29.0 CS3: Discovers active Starlark packages with
// filters.pre_completion permission and registers them in the chain.
// Called at startup and after package install/permission changes.
package filters
import (
"context"
"log"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/sandbox"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// DiscoverStarlarkFilters scans the package registry for active starlark
// extensions with granted filters.pre_completion permission and registers
// them in the filter chain. Extension filters start at order 100.
//
// Called once at startup. Future: also called on package install/activate
// to hot-register new filters without restart.
func DiscoverStarlarkFilters(ctx context.Context, chain *Chain, stores store.Stores, runner *sandbox.Runner) {
if stores.ExtPermissions == nil || stores.Packages == nil {
return
}
pkgs, err := stores.Packages.ListEnabledByType(ctx, "extension")
if err != nil {
log.Printf("⚠ filter discovery: failed to list extensions: %v", err)
return
}
// Also check "full" packages (surface + extension)
fullPkgs, err := stores.Packages.ListEnabledByType(ctx, "full")
if err == nil {
pkgs = append(pkgs, fullPkgs...)
}
registered := 0
for i := range pkgs {
pkg := &pkgs[i]
// Must be starlark tier and active
if pkg.Tier != models.ExtTierStarlark {
continue
}
if pkg.Status != models.PackageStatusActive {
continue
}
// Must have the pre_completion filter permission granted
granted, err := stores.ExtPermissions.GrantedForPackage(ctx, pkg.ID)
if err != nil {
continue
}
hasFilterPerm := false
for _, p := range granted {
if p == models.ExtPermFiltersPreCompletion {
hasFilterPerm = true
break
}
}
if !hasFilterPerm {
continue
}
// Must have a starlark script
if _, ok := pkg.Manifest["_starlark_script"].(string); !ok {
continue
}
// Register with order 100+ (extension range, after built-ins)
order := 100 + registered
chain.Register(NewStarlarkFilter(runner, pkg, order))
registered++
}
if registered > 0 {
log.Printf(" 🔗 Discovered %d Starlark pre-completion filter(s)", registered)
}
}

134
server/filters/filters.go Normal file
View File

@@ -0,0 +1,134 @@
// Package filters — pre-completion filter chain.
//
// v0.29.0: Reference implementation for the server-side filter model.
// Filters run before every completion request, injecting context
// (system messages, knowledge base results, external data) into the
// conversation. The chain is composable: Go built-in filters register
// at startup; Starlark extension filters register at package install.
//
// Architecture:
//
// user message → [filter chain] → LLM
// │
// ├─ KB auto-inject (built-in, CS0)
// ├─ memory hint (future migration)
// └─ ext filters (Starlark, CS3)
//
// Each filter receives a CompletionContext and returns zero or more
// system messages to inject. Filters are executed in registration
// order. A filter failure is logged and skipped — it never aborts
// the completion.
package filters
import (
"context"
"log"
"sort"
"time"
)
// ─── Context ─────────────────────────────────
// CompletionContext carries the state available to pre-completion filters.
// Filters should treat this as read-only.
type CompletionContext struct {
ChannelID string
UserID string
PersonaID string
LastUserMessage string // current user message content (for semantic search)
TeamIDs []string
}
// ─── Contribution ────────────────────────────
// InjectedMessage is a role + content pair that a filter wants to inject
// into the conversation. Kept simple to avoid coupling to provider types.
type InjectedMessage struct {
Role string // "system" (most common), "user", or "assistant"
Content string
}
// Contribution holds what a filter wants to inject into the completion.
type Contribution struct {
// Messages are injected into the conversation in order, after the
// persona/project system prompts and before the message history.
Messages []InjectedMessage
}
// ─── Filter Interface ────────────────────────
// PreCompletionFilter processes context before a completion request.
// Implementations must be safe for concurrent use.
type PreCompletionFilter interface {
// Name returns a human-readable identifier (e.g., "kb-auto-inject").
// Used in logs and admin UI.
Name() string
// Order returns the execution priority. Lower values run first.
// Built-in filters use 099; extension filters use 100+.
Order() int
// Execute runs the filter and returns messages to inject.
// Returning nil or an empty Contribution is valid (no injection).
// Errors are logged and skipped — they never abort the completion.
Execute(ctx context.Context, cc *CompletionContext) (*Contribution, error)
}
// ─── Chain ───────────────────────────────────
// Chain holds an ordered set of pre-completion filters.
type Chain struct {
filters []PreCompletionFilter
sorted bool
}
// NewChain creates an empty filter chain.
func NewChain() *Chain {
return &Chain{}
}
// Register adds a filter to the chain. Filters are sorted by Order()
// before the first execution.
func (c *Chain) Register(f PreCompletionFilter) {
c.filters = append(c.filters, f)
c.sorted = false
log.Printf(" 🔗 Pre-completion filter registered: %s (order=%d)", f.Name(), f.Order())
}
// Len returns the number of registered filters.
func (c *Chain) Len() int {
return len(c.filters)
}
// Execute runs all filters in order, collecting their contributions.
// Returns the combined list of messages to inject. Individual filter
// errors are logged and skipped.
func (c *Chain) Execute(ctx context.Context, cc *CompletionContext) []InjectedMessage {
if len(c.filters) == 0 {
return nil
}
if !c.sorted {
sort.Slice(c.filters, func(i, j int) bool {
return c.filters[i].Order() < c.filters[j].Order()
})
c.sorted = true
}
var messages []InjectedMessage
for _, f := range c.filters {
start := time.Now()
contrib, err := f.Execute(ctx, cc)
elapsed := time.Since(start)
if err != nil {
log.Printf("⚠ filter %s failed (%.0fms): %v", f.Name(), elapsed.Seconds()*1000, err)
continue
}
if contrib != nil && len(contrib.Messages) > 0 {
messages = append(messages, contrib.Messages...)
log.Printf(" 🔗 filter %s: injected %d messages (%.0fms)", f.Name(), len(contrib.Messages), elapsed.Seconds()*1000)
}
}
return messages
}

View File

@@ -0,0 +1,155 @@
package filters
import (
"context"
"errors"
"testing"
)
// ── Mock filter ─────────────────────────────
type mockFilter struct {
name string
order int
messages []InjectedMessage
err error
}
func (m *mockFilter) Name() string { return m.name }
func (m *mockFilter) Order() int { return m.order }
func (m *mockFilter) Execute(_ context.Context, _ *CompletionContext) (*Contribution, error) {
if m.err != nil {
return nil, m.err
}
if len(m.messages) == 0 {
return nil, nil
}
return &Contribution{Messages: m.messages}, nil
}
// ── Tests ───────────────────────────────────
func TestChainEmpty(t *testing.T) {
c := NewChain()
result := c.Execute(context.Background(), &CompletionContext{})
if len(result) != 0 {
t.Fatalf("expected 0 messages, got %d", len(result))
}
}
func TestChainSingleFilter(t *testing.T) {
c := NewChain()
c.Register(&mockFilter{
name: "test",
order: 0,
messages: []InjectedMessage{
{Role: "system", Content: "hello"},
},
})
result := c.Execute(context.Background(), &CompletionContext{
ChannelID: "ch-1",
UserID: "u-1",
})
if len(result) != 1 {
t.Fatalf("expected 1 message, got %d", len(result))
}
if result[0].Content != "hello" {
t.Fatalf("expected 'hello', got %q", result[0].Content)
}
}
func TestChainOrderRespected(t *testing.T) {
c := NewChain()
// Register out of order — chain should sort by Order()
c.Register(&mockFilter{
name: "second",
order: 20,
messages: []InjectedMessage{{Role: "system", Content: "B"}},
})
c.Register(&mockFilter{
name: "first",
order: 10,
messages: []InjectedMessage{{Role: "system", Content: "A"}},
})
result := c.Execute(context.Background(), &CompletionContext{})
if len(result) != 2 {
t.Fatalf("expected 2 messages, got %d", len(result))
}
if result[0].Content != "A" {
t.Fatalf("expected 'A' first, got %q", result[0].Content)
}
if result[1].Content != "B" {
t.Fatalf("expected 'B' second, got %q", result[1].Content)
}
}
func TestChainErrorIsolation(t *testing.T) {
c := NewChain()
c.Register(&mockFilter{
name: "failing",
order: 0,
err: errors.New("boom"),
})
c.Register(&mockFilter{
name: "surviving",
order: 10,
messages: []InjectedMessage{{Role: "system", Content: "ok"}},
})
result := c.Execute(context.Background(), &CompletionContext{})
if len(result) != 1 {
t.Fatalf("expected 1 message (failing filter skipped), got %d", len(result))
}
if result[0].Content != "ok" {
t.Fatalf("expected 'ok', got %q", result[0].Content)
}
}
func TestChainNilContribution(t *testing.T) {
c := NewChain()
c.Register(&mockFilter{
name: "noop",
order: 0,
// nil messages → nil contribution
})
result := c.Execute(context.Background(), &CompletionContext{})
if len(result) != 0 {
t.Fatalf("expected 0 messages from noop filter, got %d", len(result))
}
}
func TestChainMultipleMessages(t *testing.T) {
c := NewChain()
c.Register(&mockFilter{
name: "multi",
order: 0,
messages: []InjectedMessage{
{Role: "system", Content: "first"},
{Role: "system", Content: "second"},
},
})
result := c.Execute(context.Background(), &CompletionContext{})
if len(result) != 2 {
t.Fatalf("expected 2 messages, got %d", len(result))
}
}
func TestChainLen(t *testing.T) {
c := NewChain()
if c.Len() != 0 {
t.Fatalf("expected 0, got %d", c.Len())
}
c.Register(&mockFilter{name: "a", order: 0})
c.Register(&mockFilter{name: "b", order: 1})
if c.Len() != 2 {
t.Fatalf("expected 2, got %d", c.Len())
}
}

112
server/filters/kb_inject.go Normal file
View File

@@ -0,0 +1,112 @@
// Package filters — kb_inject.go
//
// v0.29.0 CS0: Built-in pre-completion filter that auto-injects KB
// context into the system prompt. Refactored from the inline
// BuildKBHint() function in knowledge_bases.go.
//
// This is the reference implementation for the filter chain model
// that Starlark extensions will mirror in CS3.
package filters
import (
"context"
"fmt"
"log"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// KBInjectFilter injects a system prompt listing active knowledge bases
// so the LLM knows to use the kb_search tool.
type KBInjectFilter struct {
stores store.Stores
}
// NewKBInjectFilter creates the built-in KB auto-inject filter.
func NewKBInjectFilter(stores store.Stores) *KBInjectFilter {
return &KBInjectFilter{stores: stores}
}
func (f *KBInjectFilter) Name() string { return "kb-auto-inject" }
func (f *KBInjectFilter) Order() int { return 10 } // built-in, runs early
func (f *KBInjectFilter) Execute(ctx context.Context, cc *CompletionContext) (*Contribution, error) {
type kbInfo struct {
Name string
DocCount int
}
var kbs []kbInfo
seen := make(map[string]bool)
// ── Persona-bound KBs (v0.17.0) ──
if cc.PersonaID != "" {
personaKBs, err := f.stores.Personas.GetKBs(ctx, cc.PersonaID)
if err == nil {
for _, pkb := range personaKBs {
if pkb.ChunkCount > 0 {
kbs = append(kbs, kbInfo{Name: pkb.KBName, DocCount: pkb.DocumentCount})
seen[pkb.KBID] = true
}
}
}
}
// ── Project-bound KBs (v0.19.0) ──
if f.stores.Projects != nil {
projID, _ := f.stores.Projects.GetProjectIDForChannel(ctx, cc.ChannelID)
if projID != "" {
projKBIDs, projErr := f.stores.Projects.GetKBIDs(ctx, projID)
if projErr == nil {
for _, kbID := range projKBIDs {
if !seen[kbID] {
kb, kbErr := f.stores.KnowledgeBases.GetByID(ctx, kbID)
if kbErr == nil && kb.ChunkCount > 0 {
kbs = append(kbs, kbInfo{Name: kb.Name, DocCount: kb.DocumentCount})
seen[kbID] = true
}
}
}
}
}
}
// ── Channel-linked KBs ──
channelKBs, err := f.stores.KnowledgeBases.GetChannelKBs(ctx, cc.ChannelID)
if err == nil {
for _, ckb := range channelKBs {
if ckb.Enabled && ckb.DocumentCount > 0 && !seen[ckb.KBID] {
kbs = append(kbs, kbInfo{Name: ckb.KBName, DocCount: ckb.DocumentCount})
seen[ckb.KBID] = true
}
}
}
// ── Personal KBs (always available to owner) ──
personalKBs, err := f.stores.KnowledgeBases.ListPersonal(ctx, cc.UserID)
if err == nil {
for _, kb := range personalKBs {
if !seen[kb.ID] && kb.ChunkCount > 0 {
kbs = append(kbs, kbInfo{Name: kb.Name, DocCount: kb.DocumentCount})
seen[kb.ID] = true
}
}
}
if len(kbs) == 0 {
return nil, nil
}
hint := "\nYou have access to the following knowledge bases:\n"
for _, kb := range kbs {
hint += fmt.Sprintf("- \"%s\" (%d documents)\n", kb.Name, kb.DocCount)
}
hint += "Use the kb_search tool to find relevant information from these sources when the user asks questions that might be answered by their documents."
log.Printf(" 🔗 kb-auto-inject: %d KBs active for channel %s", len(kbs), cc.ChannelID[:min(8, len(cc.ChannelID))])
return &Contribution{
Messages: []InjectedMessage{
{Role: "system", Content: hint},
},
}, nil
}

View File

@@ -0,0 +1,133 @@
// Package filters — starlark_filter.go
//
// v0.29.0 CS3: Bridges the pre-completion filter chain to Starlark
// extension scripts. When a package declares "filters.pre_completion"
// permission and is active, this filter runs the package's
// on_pre_completion(ctx) function and converts the return value
// into injected messages.
//
// Starlark contract:
//
// def on_pre_completion(ctx):
// """Called before each completion request.
// ctx is a dict: {"channel_id": "...", "user_id": "...", "persona_id": "..."}
// Return a list of dicts: [{"role": "system", "content": "..."}]
// Return None or [] to inject nothing.
// """
// return [{"role": "system", "content": "Extra context from extension"}]
package filters
import (
"context"
"fmt"
"log"
"go.starlark.net/starlark"
"git.gobha.me/xcaliber/chat-switchboard/sandbox"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// StarlarkFilter runs a Starlark package's on_pre_completion entry point.
type StarlarkFilter struct {
runner *sandbox.Runner
pkg *store.PackageRegistration
order int
}
// NewStarlarkFilter creates a filter backed by a Starlark package.
func NewStarlarkFilter(runner *sandbox.Runner, pkg *store.PackageRegistration, order int) *StarlarkFilter {
return &StarlarkFilter{
runner: runner,
pkg: pkg,
order: order,
}
}
func (f *StarlarkFilter) Name() string { return "ext:" + f.pkg.ID }
func (f *StarlarkFilter) Order() int { return f.order }
func (f *StarlarkFilter) Execute(ctx context.Context, cc *CompletionContext) (*Contribution, error) {
// Build context dict for the Starlark function
ctxDict := starlark.NewDict(4)
ctxDict.SetKey(starlark.String("channel_id"), starlark.String(cc.ChannelID))
ctxDict.SetKey(starlark.String("user_id"), starlark.String(cc.UserID))
ctxDict.SetKey(starlark.String("persona_id"), starlark.String(cc.PersonaID))
ctxDict.SetKey(starlark.String("last_user_message"), starlark.String(cc.LastUserMessage))
val, output, err := f.runner.CallEntryPoint(ctx, f.pkg, "on_pre_completion",
starlark.Tuple{ctxDict}, nil)
if err != nil {
return nil, fmt.Errorf("ext:%s: %w", f.pkg.ID, err)
}
if output != "" {
log.Printf(" 🔧 ext:%s print: %s", f.pkg.ID, output)
}
// Parse return value into messages
messages, err := parseStarlarkMessages(val)
if err != nil {
return nil, fmt.Errorf("ext:%s: invalid return value: %w", f.pkg.ID, err)
}
if len(messages) == 0 {
return nil, nil
}
return &Contribution{Messages: messages}, nil
}
// parseStarlarkMessages converts a Starlark return value to InjectedMessages.
// Accepts: None, empty list, or list of dicts with "role" and "content" keys.
func parseStarlarkMessages(val starlark.Value) ([]InjectedMessage, error) {
if val == nil || val == starlark.None {
return nil, nil
}
list, ok := val.(*starlark.List)
if !ok {
return nil, fmt.Errorf("expected list or None, got %s", val.Type())
}
if list.Len() == 0 {
return nil, nil
}
var messages []InjectedMessage
iter := list.Iterate()
defer iter.Done()
var item starlark.Value
for iter.Next(&item) {
dict, ok := item.(*starlark.Dict)
if !ok {
return nil, fmt.Errorf("expected dict in list, got %s", item.Type())
}
roleVal, found, _ := dict.Get(starlark.String("role"))
if !found {
return nil, fmt.Errorf("message dict missing 'role' key")
}
role, ok := starlark.AsString(roleVal)
if !ok {
return nil, fmt.Errorf("'role' must be a string")
}
contentVal, found, _ := dict.Get(starlark.String("content"))
if !found {
return nil, fmt.Errorf("message dict missing 'content' key")
}
content, ok := starlark.AsString(contentVal)
if !ok {
return nil, fmt.Errorf("'content' must be a string")
}
messages = append(messages, InjectedMessage{
Role: role,
Content: content,
})
}
return messages, nil
}

View File

@@ -0,0 +1,119 @@
package filters
import (
"testing"
"go.starlark.net/starlark"
)
func TestParseStarlarkMessages_None(t *testing.T) {
msgs, err := parseStarlarkMessages(starlark.None)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if msgs != nil {
t.Fatalf("expected nil, got %d messages", len(msgs))
}
}
func TestParseStarlarkMessages_Nil(t *testing.T) {
msgs, err := parseStarlarkMessages(nil)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if msgs != nil {
t.Fatalf("expected nil, got %d messages", len(msgs))
}
}
func TestParseStarlarkMessages_EmptyList(t *testing.T) {
msgs, err := parseStarlarkMessages(starlark.NewList(nil))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if msgs != nil {
t.Fatalf("expected nil, got %d messages", len(msgs))
}
}
func TestParseStarlarkMessages_SingleMessage(t *testing.T) {
d := starlark.NewDict(2)
d.SetKey(starlark.String("role"), starlark.String("system"))
d.SetKey(starlark.String("content"), starlark.String("injected context"))
list := starlark.NewList([]starlark.Value{d})
msgs, err := parseStarlarkMessages(list)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(msgs) != 1 {
t.Fatalf("expected 1 message, got %d", len(msgs))
}
if msgs[0].Role != "system" {
t.Fatalf("expected role 'system', got %q", msgs[0].Role)
}
if msgs[0].Content != "injected context" {
t.Fatalf("expected content 'injected context', got %q", msgs[0].Content)
}
}
func TestParseStarlarkMessages_MultipleMessages(t *testing.T) {
d1 := starlark.NewDict(2)
d1.SetKey(starlark.String("role"), starlark.String("system"))
d1.SetKey(starlark.String("content"), starlark.String("first"))
d2 := starlark.NewDict(2)
d2.SetKey(starlark.String("role"), starlark.String("system"))
d2.SetKey(starlark.String("content"), starlark.String("second"))
list := starlark.NewList([]starlark.Value{d1, d2})
msgs, err := parseStarlarkMessages(list)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(msgs) != 2 {
t.Fatalf("expected 2 messages, got %d", len(msgs))
}
if msgs[0].Content != "first" || msgs[1].Content != "second" {
t.Fatalf("wrong content: %q, %q", msgs[0].Content, msgs[1].Content)
}
}
func TestParseStarlarkMessages_MissingRole(t *testing.T) {
d := starlark.NewDict(1)
d.SetKey(starlark.String("content"), starlark.String("no role"))
list := starlark.NewList([]starlark.Value{d})
_, err := parseStarlarkMessages(list)
if err == nil {
t.Fatal("expected error for missing role")
}
}
func TestParseStarlarkMessages_MissingContent(t *testing.T) {
d := starlark.NewDict(1)
d.SetKey(starlark.String("role"), starlark.String("system"))
list := starlark.NewList([]starlark.Value{d})
_, err := parseStarlarkMessages(list)
if err == nil {
t.Fatal("expected error for missing content")
}
}
func TestParseStarlarkMessages_WrongType(t *testing.T) {
_, err := parseStarlarkMessages(starlark.String("not a list"))
if err == nil {
t.Fatal("expected error for wrong type")
}
}
func TestParseStarlarkMessages_NonDictInList(t *testing.T) {
list := starlark.NewList([]starlark.Value{starlark.String("not a dict")})
_, err := parseStarlarkMessages(list)
if err == nil {
t.Fatal("expected error for non-dict in list")
}
}

View File

@@ -126,10 +126,7 @@ func (h *AdminHandler) UpdateUserRole(c *gin.Context) {
return return
} }
if user.Role == models.UserRoleAdmin { if user.Role == models.UserRoleAdmin {
var adminCount int adminCount, _ := h.stores.Users.CountByRole(c.Request.Context(), models.UserRoleAdmin)
database.DB.QueryRowContext(c.Request.Context(),
database.Q(`SELECT COUNT(*) FROM users WHERE role = $1`),
models.UserRoleAdmin).Scan(&adminCount)
if adminCount <= 1 { if adminCount <= 1 {
c.JSON(http.StatusConflict, gin.H{"error": "cannot demote the last admin"}) c.JSON(http.StatusConflict, gin.H{"error": "cannot demote the last admin"})
return return
@@ -203,7 +200,7 @@ func (h *AdminHandler) destroyVault(c *gin.Context, userID string) {
return return
} }
deleted := DestroyVaultDB(c.Request.Context(), userID) deleted := DestroyVaultDB(c.Request.Context(), h.stores, userID)
// Evict from session cache (if user is currently logged in) // Evict from session cache (if user is currently logged in)
h.uekCache.Evict(userID) h.uekCache.Evict(userID)
@@ -236,7 +233,7 @@ func (h *AdminHandler) ResetVault(c *gin.Context) {
return return
} }
deleted := DestroyVaultDB(c.Request.Context(), userID) deleted := DestroyVaultDB(c.Request.Context(), h.stores, userID)
h.uekCache.Evict(userID) h.uekCache.Evict(userID)
h.auditLog(c, "user.vault_reset", "user", userID, models.JSONMap{ h.auditLog(c, "user.vault_reset", "user", userID, models.JSONMap{
@@ -261,10 +258,7 @@ func (h *AdminHandler) DeleteUser(c *gin.Context) {
return return
} }
if user.Role == models.UserRoleAdmin { if user.Role == models.UserRoleAdmin {
var adminCount int adminCount, _ := h.stores.Users.CountByRole(c.Request.Context(), models.UserRoleAdmin)
database.DB.QueryRowContext(c.Request.Context(),
database.Q(`SELECT COUNT(*) FROM users WHERE role = $1`),
models.UserRoleAdmin).Scan(&adminCount)
if adminCount <= 1 { if adminCount <= 1 {
c.JSON(http.StatusConflict, gin.H{"error": "cannot delete the last admin"}) c.JSON(http.StatusConflict, gin.H{"error": "cannot delete the last admin"})
return return
@@ -769,10 +763,9 @@ func (h *AdminHandler) GetStats(c *gin.Context) {
ctx := c.Request.Context() ctx := c.Request.Context()
stats := gin.H{} stats := gin.H{}
var userCount, channelCount, messageCount int userCount, _ := h.stores.Users.CountAll(ctx)
database.DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&userCount) channelCount, _ := h.stores.Channels.CountAll(ctx)
database.DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM channels").Scan(&channelCount) messageCount, _ := h.stores.Messages.CountAll(ctx)
database.DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM messages").Scan(&messageCount)
stats["users"] = userCount stats["users"] = userCount
stats["channels"] = channelCount stats["channels"] = channelCount

View File

@@ -20,7 +20,6 @@ import (
"git.gobha.me/xcaliber/chat-switchboard/auth" "git.gobha.me/xcaliber/chat-switchboard/auth"
"git.gobha.me/xcaliber/chat-switchboard/config" "git.gobha.me/xcaliber/chat-switchboard/config"
"git.gobha.me/xcaliber/chat-switchboard/crypto" "git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
@@ -184,10 +183,7 @@ func (h *AuthHandler) OIDCLogin(c *gin.Context) {
} }
// Store state for callback verification // Store state for callback verification
_, err := database.DB.ExecContext(c.Request.Context(), database.Q(` if err := h.stores.GlobalConfig.SaveOIDCState(c.Request.Context(), state, nonce, c.Query("redirect")); err != nil {
INSERT INTO oidc_auth_state (state, nonce, redirect_to) VALUES ($1, $2, $3)
`), state, nonce, c.Query("redirect"))
if err != nil {
log.Printf("[auth/oidc] warn: could not store state: %v", err) log.Printf("[auth/oidc] warn: could not store state: %v", err)
} }
@@ -220,29 +216,15 @@ func (h *AuthHandler) OIDCCallback(c *gin.Context) {
return return
} }
// Verify state // Verify state (nonce + redirectTo retrieved but not yet validated — TODO)
var nonce, redirectTo string _, _, err := h.stores.GlobalConfig.ConsumeOIDCState(c.Request.Context(), state)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT nonce, COALESCE(redirect_to, '') FROM oidc_auth_state WHERE state = $1
`), state).Scan(&nonce, &redirectTo)
if err != nil { if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid or expired state"}) c.JSON(http.StatusBadRequest, gin.H{"error": "invalid or expired state"})
return return
} }
// Clean up state (one-time use) // Clean up stale states (> 10 minutes old)
database.DB.ExecContext(c.Request.Context(), database.Q(` _ = h.stores.GlobalConfig.CleanupOIDCState(c.Request.Context())
DELETE FROM oidc_auth_state WHERE state = $1
`), state)
// Also clean up stale states (> 10 minutes old)
if database.IsSQLite() {
database.DB.ExecContext(c.Request.Context(),
`DELETE FROM oidc_auth_state WHERE created_at < datetime('now', '-10 minutes')`)
} else {
database.DB.ExecContext(c.Request.Context(),
`DELETE FROM oidc_auth_state WHERE created_at < NOW() - INTERVAL '10 minutes'`)
}
// Determine redirect URI (must match what was sent in the login request) // Determine redirect URI (must match what was sent in the login request)
redirectURI := h.cfg.OIDCRedirectURL redirectURI := h.cfg.OIDCRedirectURL
@@ -363,24 +345,17 @@ func hashToken(token string) string {
// - AdminHandler.destroyVault (admin-initiated reset) // - AdminHandler.destroyVault (admin-initiated reset)
// //
// Does NOT evict from UEK cache or write audit logs — callers handle that. // Does NOT evict from UEK cache or write audit logs — callers handle that.
func DestroyVaultDB(ctx context.Context, userID string) (providersDeleted int64) { // v0.29.0: accepts stores instead of using database.DB directly.
_, err := database.DB.ExecContext(ctx, database.Q(` func DestroyVaultDB(ctx context.Context, stores store.Stores, userID string) (providersDeleted int64) {
UPDATE users if err := stores.Users.ClearVaultKeys(ctx, userID); err != nil {
SET encrypted_uek = NULL, uek_salt = NULL, uek_nonce = NULL, vault_set = false
WHERE id = $1
`), userID)
if err != nil {
log.Printf("⚠ DestroyVaultDB: failed to clear vault columns for user %s: %v", userID, err) log.Printf("⚠ DestroyVaultDB: failed to clear vault columns for user %s: %v", userID, err)
} }
result, err := database.DB.ExecContext(ctx, database.Q(` rows, err := stores.Providers.DeletePersonalByOwner(ctx, userID)
DELETE FROM provider_configs WHERE scope = 'personal' AND owner_id = $1
`), userID)
if err != nil { if err != nil {
log.Printf("⚠ DestroyVaultDB: failed to delete personal providers for user %s: %v", userID, err) log.Printf("⚠ DestroyVaultDB: failed to delete personal providers for user %s: %v", userID, err)
return 0 return 0
} }
rows, _ := result.RowsAffected()
return rows return rows
} }
@@ -392,14 +367,9 @@ func DestroyVaultDB(ctx context.Context, userID string) (providersDeleted int64)
// //
// Used by BootstrapAdmin and SeedUsers where the password is known at // Used by BootstrapAdmin and SeedUsers where the password is known at
// startup but the UEK cache is not available. // startup but the UEK cache is not available.
func ProbeAndRepairVault(ctx context.Context, userID, password string) { // v0.29.0: accepts stores instead of using database.DB directly.
var vaultSet bool func ProbeAndRepairVault(ctx context.Context, stores store.Stores, userID, password string) {
var encryptedUEK, salt, nonce []byte vaultSet, encryptedUEK, salt, nonce, err := stores.Users.GetVaultKeys(ctx, userID)
err := database.DB.QueryRowContext(ctx, database.Q(`
SELECT vault_set, encrypted_uek, uek_salt, uek_nonce
FROM users WHERE id = $1
`), userID).Scan(&vaultSet, &encryptedUEK, &salt, &nonce)
if err != nil || !vaultSet { if err != nil || !vaultSet {
return // no vault to probe return // no vault to probe
} }
@@ -410,7 +380,7 @@ func ProbeAndRepairVault(ctx context.Context, userID, password string) {
} }
// Stale seal: password has actually changed since the vault was sealed // Stale seal: password has actually changed since the vault was sealed
deleted := DestroyVaultDB(ctx, userID) deleted := DestroyVaultDB(ctx, stores, userID)
if deleted > 0 { if deleted > 0 {
log.Printf(" 🔐 Vault stale-seal repair for user %s: cleared %d personal provider(s)", userID, deleted) log.Printf(" 🔐 Vault stale-seal repair for user %s: cleared %d personal provider(s)", userID, deleted)
} else { } else {
@@ -442,12 +412,7 @@ func (h *AuthHandler) initVault(ctx context.Context, userID, password string) er
return err return err
} }
_, err = database.DB.ExecContext(ctx, database.Q(` if err := h.stores.Users.InitVaultKeys(ctx, userID, encryptedUEK, salt, nonce); err != nil {
UPDATE users
SET encrypted_uek = $1, uek_salt = $2, uek_nonce = $3, vault_set = true
WHERE id = $4
`), encryptedUEK, salt, nonce, userID)
if err != nil {
return err return err
} }
@@ -464,13 +429,7 @@ func (h *AuthHandler) unlockVault(ctx context.Context, user *models.User, passwo
return return
} }
var vaultSet bool vaultSet, encryptedUEK, salt, nonce, err := h.stores.Users.GetVaultKeys(ctx, user.ID)
var encryptedUEK, salt, nonce []byte
err := database.DB.QueryRowContext(ctx, database.Q(`
SELECT vault_set, encrypted_uek, uek_salt, uek_nonce
FROM users WHERE id = $1
`), user.ID).Scan(&vaultSet, &encryptedUEK, &salt, &nonce)
if err != nil { if err != nil {
log.Printf("⚠ Vault query failed for user %s: %v", user.ID, err) log.Printf("⚠ Vault query failed for user %s: %v", user.ID, err)
return return
@@ -489,10 +448,7 @@ func (h *AuthHandler) unlockVault(ctx context.Context, user *models.User, passwo
uek, err := crypto.UnwrapUEK(encryptedUEK, nonce, pdk) uek, err := crypto.UnwrapUEK(encryptedUEK, nonce, pdk)
if err != nil { if err != nil {
// Stale seal: encrypted_uek was wrapped with a different password // Stale seal: encrypted_uek was wrapped with a different password
// (e.g. admin reset, BootstrapAdmin rotation, or pre-UEK migration). deleted := DestroyVaultDB(ctx, h.stores, user.ID)
// The old UEK is irrecoverable. Destroy the vault and re-initialize
// with the current password so the user isn't permanently locked out.
deleted := DestroyVaultDB(ctx, user.ID)
if deleted > 0 { if deleted > 0 {
log.Printf("⚠ Vault stale-seal recovery for user %s: cleared %d personal provider(s)", user.ID, deleted) log.Printf("⚠ Vault stale-seal recovery for user %s: cleared %d personal provider(s)", user.ID, deleted)
} else { } else {
@@ -536,7 +492,7 @@ func BootstrapAdmin(cfg *config.Config, s store.Stores) {
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(cfg.AdminUsername)) handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(cfg.AdminUsername))
s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle}) s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle})
} }
ProbeAndRepairVault(ctx, existing.ID, cfg.AdminPassword) ProbeAndRepairVault(ctx, s, existing.ID, cfg.AdminPassword)
log.Printf(" ✅ Admin user '%s' updated", cfg.AdminUsername) log.Printf(" ✅ Admin user '%s' updated", cfg.AdminUsername)
return return
} }
@@ -626,7 +582,7 @@ func SeedUsers(cfg *config.Config, s store.Stores) {
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(username)) handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(username))
s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle}) s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle})
} }
ProbeAndRepairVault(ctx, existing.ID, password) ProbeAndRepairVault(ctx, s, existing.ID, password)
log.Printf(" 🌱 Seed user '%s' updated (role=%s)", username, role) log.Printf(" 🌱 Seed user '%s' updated (role=%s)", username, role)
continue continue
} }

View File

@@ -15,7 +15,8 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
const avatarSize = 128 const avatarSize = 128
@@ -76,10 +77,7 @@ func (h *SettingsHandler) UploadAvatar(c *gin.Context) {
dataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(buf.Bytes()) dataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(buf.Bytes())
// Store in DB // Store in DB
_, err = database.DB.Exec( err = h.stores.Users.Update(c.Request.Context(), userID, map[string]interface{}{"avatar_url": dataURI})
database.Q(`UPDATE users SET avatar_url = $1, updated_at = NOW() WHERE id = $2`),
dataURI, userID,
)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to save avatar"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to save avatar"})
return return
@@ -93,10 +91,7 @@ func (h *SettingsHandler) UploadAvatar(c *gin.Context) {
func (h *SettingsHandler) DeleteAvatar(c *gin.Context) { func (h *SettingsHandler) DeleteAvatar(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
_, err := database.DB.Exec( err := h.stores.Users.Update(c.Request.Context(), userID, map[string]interface{}{"avatar_url": nil})
database.Q(`UPDATE users SET avatar_url = NULL, updated_at = NOW() WHERE id = $1`),
userID,
)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to remove avatar"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to remove avatar"})
return return
@@ -166,7 +161,9 @@ func bilinearMix(c00, c10, c01, c11 color.Color, xf, yf float64) color.Color {
// ── Persona Avatar Upload ──────────────────── // ── Persona Avatar Upload ────────────────────
// POST /api/v1/personas/:id/avatar (user) or /api/v1/admin/personas/:id/avatar (admin) // POST /api/v1/personas/:id/avatar (user) or /api/v1/admin/personas/:id/avatar (admin)
func UploadPersonaAvatar(c *gin.Context) { // UploadPersonaAvatar uploads and resizes a persona avatar.
// v0.29.0: accepts PersonaStore instead of using database.DB directly.
func UploadPersonaAvatar(personas store.PersonaStore, c *gin.Context) {
personaID := c.Param("id") personaID := c.Param("id")
var req uploadAvatarRequest var req uploadAvatarRequest
@@ -206,40 +203,26 @@ func UploadPersonaAvatar(c *gin.Context) {
dataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(buf.Bytes()) dataURI := "data:image/png;base64," + base64.StdEncoding.EncodeToString(buf.Bytes())
result, err := database.DB.Exec( err = personas.Update(c.Request.Context(), personaID, models.PersonaPatch{Avatar: &dataURI})
database.Q(`UPDATE personas SET avatar = $1, updated_at = NOW() WHERE id = $2`),
dataURI, personaID,
)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to save avatar"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to save avatar"})
return return
} }
rows, _ := result.RowsAffected()
if rows == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "persona not found"})
return
}
c.JSON(http.StatusOK, gin.H{"avatar": dataURI}) c.JSON(http.StatusOK, gin.H{"avatar": dataURI})
} }
// DeletePersonaAvatar clears a persona's avatar. // DeletePersonaAvatar clears a persona's avatar.
func DeletePersonaAvatar(c *gin.Context) { // v0.29.0: accepts PersonaStore instead of using database.DB directly.
func DeletePersonaAvatar(personas store.PersonaStore, c *gin.Context) {
personaID := c.Param("id") personaID := c.Param("id")
result, err := database.DB.Exec( empty := ""
database.Q(`UPDATE personas SET avatar = '', updated_at = NOW() WHERE id = $1`), err := personas.Update(c.Request.Context(), personaID, models.PersonaPatch{Avatar: &empty})
personaID,
)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to remove avatar"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to remove avatar"})
return return
} }
rows, _ := result.RowsAffected()
if rows == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "persona not found"})
return
}
c.JSON(http.StatusOK, gin.H{"message": "avatar removed"}) c.JSON(http.StatusOK, gin.H{"message": "avatar removed"})
} }

View File

@@ -10,7 +10,6 @@ import (
capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities" capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities"
"git.gobha.me/xcaliber/chat-switchboard/auth" "git.gobha.me/xcaliber/chat-switchboard/auth"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/health" "git.gobha.me/xcaliber/chat-switchboard/health"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
@@ -90,10 +89,11 @@ func (h *ModelHandler) buildHealthMap(ctx context.Context) map[string]models.Pro
} }
// ResolveModelCaps is the canonical capability resolver for any model. // ResolveModelCaps is the canonical capability resolver for any model.
func ResolveModelCaps(c *gin.Context, modelID, configID string) models.ModelCapabilities { // v0.29.0: accepts CatalogStore instead of using database.DB directly.
func ResolveModelCaps(catalog store.CatalogStore, modelID, configID string) models.ModelCapabilities {
// 1. Exact match: model_id + provider_config_id // 1. Exact match: model_id + provider_config_id
if configID != "" { if configID != "" {
caps, ok := capsFromCatalog(modelID, configID) caps, ok := capsFromCatalog(catalog, modelID, configID)
if ok { if ok {
caps.MaxOutputTokens = capspkg.ResolveMaxOutput(modelID, caps) caps.MaxOutputTokens = capspkg.ResolveMaxOutput(modelID, caps)
return caps return caps
@@ -101,7 +101,7 @@ func ResolveModelCaps(c *gin.Context, modelID, configID string) models.ModelCapa
} }
// 2. Any provider: same model_id, any config // 2. Any provider: same model_id, any config
caps, ok := capsFromCatalog(modelID, "") caps, ok := capsFromCatalog(catalog, modelID, "")
if ok { if ok {
caps.MaxOutputTokens = capspkg.ResolveMaxOutput(modelID, caps) caps.MaxOutputTokens = capspkg.ResolveMaxOutput(modelID, caps)
return caps return caps
@@ -114,23 +114,17 @@ func ResolveModelCaps(c *gin.Context, modelID, configID string) models.ModelCapa
} }
// capsFromCatalog looks up capabilities from the model_catalog table. // capsFromCatalog looks up capabilities from the model_catalog table.
func capsFromCatalog(modelID, configID string) (models.ModelCapabilities, bool) { func capsFromCatalog(catalog store.CatalogStore, modelID, configID string) (models.ModelCapabilities, bool) {
if database.DB == nil { if catalog == nil {
return models.ModelCapabilities{}, false return models.ModelCapabilities{}, false
} }
var capsJSON []byte var capsJSON []byte
var err error var err error
if configID != "" { if configID != "" {
err = database.DB.QueryRow(database.Q(` capsJSON, err = catalog.GetCapabilities(context.Background(), modelID, configID)
SELECT capabilities FROM model_catalog
WHERE model_id = $1 AND provider_config_id = $2
`), modelID, configID).Scan(&capsJSON)
} else { } else {
err = database.DB.QueryRow(database.Q(` capsJSON, err = catalog.GetCapabilitiesAny(context.Background(), modelID)
SELECT capabilities FROM model_catalog
WHERE model_id = $1 ORDER BY last_synced_at DESC LIMIT 1
`), modelID).Scan(&capsJSON)
} }
if err != nil || len(capsJSON) == 0 { if err != nil || len(capsJSON) == 0 {
return models.ModelCapabilities{}, false return models.ModelCapabilities{}, false

View File

@@ -1,7 +1,6 @@
package handlers package handlers
import ( import (
"database/sql"
"encoding/json" "encoding/json"
"math" "math"
"net/http" "net/http"
@@ -9,69 +8,68 @@ import (
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/lib/pq"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
// ── Request / Response types ──────────────── // ── Request / Response types ────────────────
type createChannelRequest struct { type createChannelRequest struct {
Title string `json:"title" binding:"required,max=500"` Title string `json:"title" binding:"required,max=500"`
Type string `json:"type,omitempty"` // direct (default), dm, group, channel, workflow Type string `json:"type,omitempty"` // direct (default), dm, group, channel, workflow
Description string `json:"description,omitempty"` Description string `json:"description,omitempty"`
Model string `json:"model,omitempty"` Model string `json:"model,omitempty"`
SystemPrompt string `json:"system_prompt,omitempty"` SystemPrompt string `json:"system_prompt,omitempty"`
ProviderConfigID *string `json:"provider_config_id,omitempty"` ProviderConfigID *string `json:"provider_config_id,omitempty"`
Folder string `json:"folder,omitempty"` Folder string `json:"folder,omitempty"`
FolderID *string `json:"folder_id,omitempty"` FolderID *string `json:"folder_id,omitempty"`
Tags []string `json:"tags,omitempty"` Tags []string `json:"tags,omitempty"`
// v0.23.2: DM creation // v0.23.2: DM creation
Participants []string `json:"participants,omitempty"` // user IDs for DM (exactly 1 other user) Participants []string `json:"participants,omitempty"` // user IDs for DM (exactly 1 other user)
AiMode string `json:"ai_mode,omitempty"` // auto (default), mention_only, off AiMode string `json:"ai_mode,omitempty"` // auto (default), mention_only, off
} }
type updateChannelRequest struct { type updateChannelRequest struct {
Title *string `json:"title,omitempty"` Title *string `json:"title,omitempty"`
Description *string `json:"description,omitempty"` Description *string `json:"description,omitempty"`
Model *string `json:"model,omitempty"` Model *string `json:"model,omitempty"`
SystemPrompt *string `json:"system_prompt,omitempty"` SystemPrompt *string `json:"system_prompt,omitempty"`
ProviderConfigID *string `json:"provider_config_id,omitempty"` ProviderConfigID *string `json:"provider_config_id,omitempty"`
IsArchived *bool `json:"is_archived,omitempty"` IsArchived *bool `json:"is_archived,omitempty"`
IsPinned *bool `json:"is_pinned,omitempty"` IsPinned *bool `json:"is_pinned,omitempty"`
Folder *string `json:"folder,omitempty"` Folder *string `json:"folder,omitempty"`
FolderID *string `json:"folder_id,omitempty"` FolderID *string `json:"folder_id,omitempty"`
Tags []string `json:"tags,omitempty"` Tags []string `json:"tags,omitempty"`
Settings *json.RawMessage `json:"settings,omitempty"` // JSONB merge into existing settings Settings *json.RawMessage `json:"settings,omitempty"` // JSONB merge into existing settings
WorkspaceID *string `json:"workspace_id,omitempty"` // bind workspace (v0.21.5) WorkspaceID *string `json:"workspace_id,omitempty"` // bind workspace (v0.21.5)
AiMode *string `json:"ai_mode,omitempty"` // v0.23.2: auto, mention_only, off AiMode *string `json:"ai_mode,omitempty"` // v0.23.2: auto, mention_only, off
Topic *string `json:"topic,omitempty"` // v0.23.2: channel topic Topic *string `json:"topic,omitempty"` // v0.23.2: channel topic
} }
type channelResponse struct { type channelResponse struct {
ID string `json:"id"` ID string `json:"id"`
UserID string `json:"user_id"` UserID string `json:"user_id"`
Title string `json:"title"` Title string `json:"title"`
Type string `json:"type"` Type string `json:"type"`
AiMode string `json:"ai_mode,omitempty"` AiMode string `json:"ai_mode,omitempty"`
Topic *string `json:"topic,omitempty"` Topic *string `json:"topic,omitempty"`
Description *string `json:"description"` Description *string `json:"description"`
Model *string `json:"model"` Model *string `json:"model"`
ProviderConfigID *string `json:"provider_config_id"` ProviderConfigID *string `json:"provider_config_id"`
SystemPrompt *string `json:"system_prompt"` SystemPrompt *string `json:"system_prompt"`
IsArchived bool `json:"is_archived"` IsArchived bool `json:"is_archived"`
IsPinned bool `json:"is_pinned"` IsPinned bool `json:"is_pinned"`
Folder *string `json:"folder"` Folder *string `json:"folder"`
FolderID *string `json:"folder_id,omitempty"` FolderID *string `json:"folder_id,omitempty"`
ProjectID *string `json:"project_id,omitempty"` ProjectID *string `json:"project_id,omitempty"`
WorkspaceID *string `json:"workspace_id,omitempty"` WorkspaceID *string `json:"workspace_id,omitempty"`
Tags []string `json:"tags"` Tags []string `json:"tags"`
Settings json.RawMessage `json:"settings,omitempty"` Settings json.RawMessage `json:"settings,omitempty"`
MessageCount int `json:"message_count"` MessageCount int `json:"message_count"`
UnreadCount int `json:"unread_count,omitempty"` UnreadCount int `json:"unread_count,omitempty"`
CreatedAt string `json:"created_at"` CreatedAt string `json:"created_at"`
UpdatedAt string `json:"updated_at"` UpdatedAt string `json:"updated_at"`
} }
type paginatedResponse struct { type paginatedResponse struct {
@@ -83,11 +81,13 @@ type paginatedResponse struct {
} }
// ChannelHandler holds dependencies for channel endpoints. // ChannelHandler holds dependencies for channel endpoints.
type ChannelHandler struct{} type ChannelHandler struct {
stores store.Stores
}
// NewChannelHandler creates a new channel handler. // NewChannelHandler creates a new channel handler.
func NewChannelHandler() *ChannelHandler { func NewChannelHandler(stores store.Stores) *ChannelHandler {
return &ChannelHandler{} return &ChannelHandler{stores: stores}
} }
// channelDeleteHook is called after a channel is successfully deleted. // channelDeleteHook is called after a channel is successfully deleted.
@@ -100,20 +100,6 @@ func SetChannelDeleteHook(fn func(channelID string)) {
channelDeleteHook = fn channelDeleteHook = fn
} }
// ── Tag Helpers ─────────────────────────────
// writeTagsArg returns a value suitable for inserting/updating the tags column.
// Postgres: pq.Array SQLite: JSON string
func writeTagsArg(tags []string) interface{} {
if database.IsSQLite() {
b, _ := json.Marshal(tags)
return string(b)
}
return pq.Array(tags)
}
// scanJSON, scanTags, SafeJSON → safe_json.go
// ── Helpers ───────────────────────────────── // ── Helpers ─────────────────────────────────
// getUserID extracts the authenticated user's ID from context. // getUserID extracts the authenticated user's ID from context.
@@ -129,8 +115,7 @@ func isSessionAuth(c *gin.Context) bool {
} }
// sessionCanAccessChannel validates that a session participant is authorized // sessionCanAccessChannel validates that a session participant is authorized
// for the given channel. The AuthOrSession middleware already verifies this, // for the given channel.
// but handlers call this as a defense-in-depth check.
func sessionCanAccessChannel(c *gin.Context, channelID string) bool { func sessionCanAccessChannel(c *gin.Context, channelID string) bool {
if !isSessionAuth(c) { if !isSessionAuth(c) {
return false return false
@@ -156,6 +141,38 @@ func parsePagination(c *gin.Context) (page, perPage, offset int) {
return return
} }
// listItemToResponse converts a store.ChannelListItem to a handler response.
func listItemToResponse(item store.ChannelListItem) channelResponse {
tags := item.Tags
if tags == nil {
tags = []string{}
}
return channelResponse{
ID: item.ID,
UserID: item.UserID,
Title: item.Title,
Type: item.Type,
AiMode: item.AiMode,
Topic: item.Topic,
Description: item.Description,
Model: item.Model,
ProviderConfigID: item.ProviderConfigID,
SystemPrompt: item.SystemPrompt,
IsArchived: item.IsArchived,
IsPinned: item.IsPinned,
Folder: item.Folder,
FolderID: item.FolderID,
ProjectID: item.ProjectID,
WorkspaceID: item.WorkspaceID,
Tags: tags,
Settings: item.Settings,
MessageCount: item.MessageCount,
UnreadCount: item.UnreadCount,
CreatedAt: item.CreatedAt,
UpdatedAt: item.UpdatedAt,
}
}
// ── List Channels ─────────────────────────── // ── List Channels ───────────────────────────
func (h *ChannelHandler) ListChannels(c *gin.Context) { func (h *ChannelHandler) ListChannels(c *gin.Context) {
@@ -164,11 +181,7 @@ func (h *ChannelHandler) ListChannels(c *gin.Context) {
// Optional filters // Optional filters
archived := c.DefaultQuery("archived", "false") archived := c.DefaultQuery("archived", "false")
folder := c.Query("folder")
folderID := c.Query("folder_id") // UUID — preferred over folder (text)
channelType := c.DefaultQuery("type", "") // empty = all types
// v0.23.1: multi-value type filter. Supports both ?types=dm&types=channel
// and comma-joined ?types=dm,channel
var channelTypes []string var channelTypes []string
if raw := c.Query("types"); raw != "" { if raw := c.Query("types"); raw != "" {
for _, t := range strings.Split(raw, ",") { for _, t := range strings.Split(raw, ",") {
@@ -178,156 +191,31 @@ func (h *ChannelHandler) ListChannels(c *gin.Context) {
} }
} }
} }
search := strings.TrimSpace(c.Query("search")) if len(channelTypes) == 0 {
projectFilter := c.Query("project_id") // "uuid" or "none" if ct := c.DefaultQuery("type", ""); ct != "" {
channelTypes = []string{ct}
// Count total — include channels owned by user OR where user is a participant
countQuery := `SELECT COUNT(*) FROM channels c WHERE (c.user_id = $1 OR c.id IN (
SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = $1
)) AND c.is_archived = $2`
countArgs := []interface{}{userID, archived == "true"}
argN := 3
if len(channelTypes) > 0 {
placeholders := make([]string, len(channelTypes))
for i, t := range channelTypes {
placeholders[i] = "$" + strconv.Itoa(argN)
countArgs = append(countArgs, strings.TrimSpace(t))
argN++
} }
countQuery += " AND c.type IN (" + strings.Join(placeholders, ",") + ")"
} else if channelType != "" {
countQuery += ` AND c.type = $` + strconv.Itoa(argN)
countArgs = append(countArgs, channelType)
argN++
}
if folder != "" {
countQuery += ` AND c.folder = $` + strconv.Itoa(argN)
countArgs = append(countArgs, folder)
argN++
}
if folderID != "" {
countQuery += ` AND c.folder_id = $` + strconv.Itoa(argN)
countArgs = append(countArgs, folderID)
argN++
}
if search != "" {
countQuery += ` AND c.title ILIKE $` + strconv.Itoa(argN)
countArgs = append(countArgs, "%"+search+"%")
argN++
}
if projectFilter == "none" {
countQuery += ` AND c.project_id IS NULL`
} else if projectFilter != "" {
countQuery += ` AND c.project_id = $` + strconv.Itoa(argN)
countArgs = append(countArgs, projectFilter)
argN++
} }
var total int filter := store.ChannelListFilter{
if err := database.QueryRow(countQuery, countArgs...).Scan(&total); err != nil { ListOptions: store.ListOptions{Limit: perPage, Offset: offset},
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to count channels"}) Archived: archived == "true",
return Types: channelTypes,
Folder: c.Query("folder"),
FolderID: c.Query("folder_id"),
Search: strings.TrimSpace(c.Query("search")),
ProjectID: c.Query("project_id"),
} }
// Fetch channels with message count items, total, err := h.stores.Channels.ListFiltered(c.Request.Context(), userID, filter)
// Include channels owned by user OR where user is a participant (multi-user)
query := `
SELECT c.id, c.user_id, c.title, c.type, c.ai_mode, c.topic,
c.description, c.model, c.provider_config_id,
c.system_prompt, c.is_archived, c.is_pinned, c.folder, c.folder_id, c.project_id, c.workspace_id,
c.tags, c.settings,
COALESCE(mc.cnt, 0) AS message_count,
c.created_at, c.updated_at
FROM channels c
LEFT JOIN (
SELECT channel_id, COUNT(*) AS cnt FROM messages GROUP BY channel_id
) mc ON mc.channel_id = c.id
WHERE (c.user_id = $1 OR c.id IN (
SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = $1
)) AND c.is_archived = $2`
args := []interface{}{userID, archived == "true"}
argN = 3
if len(channelTypes) > 0 {
placeholders := make([]string, len(channelTypes))
for i, t := range channelTypes {
placeholders[i] = "$" + strconv.Itoa(argN)
args = append(args, strings.TrimSpace(t))
argN++
}
query += " AND c.type IN (" + strings.Join(placeholders, ",") + ")"
} else if channelType != "" {
query += ` AND c.type = $` + strconv.Itoa(argN)
args = append(args, channelType)
argN++
}
if folder != "" {
query += ` AND c.folder = $` + strconv.Itoa(argN)
args = append(args, folder)
argN++
}
if folderID != "" {
query += ` AND c.folder_id = $` + strconv.Itoa(argN)
args = append(args, folderID)
argN++
}
if search != "" {
query += ` AND c.title ILIKE $` + strconv.Itoa(argN)
args = append(args, "%"+search+"%")
argN++
}
if projectFilter == "none" {
query += ` AND c.project_id IS NULL`
} else if projectFilter != "" {
query += ` AND c.project_id = $` + strconv.Itoa(argN)
args = append(args, projectFilter)
argN++
}
query += ` ORDER BY c.is_pinned DESC, c.updated_at DESC LIMIT $` + strconv.Itoa(argN) + ` OFFSET $` + strconv.Itoa(argN+1)
args = append(args, perPage, offset)
rows, err := database.Query(query, args...)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list channels"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list channels"})
return return
} }
defer rows.Close()
channels := make([]channelResponse, 0) channels := make([]channelResponse, 0, len(items))
for rows.Next() { for _, item := range items {
var ch channelResponse channels = append(channels, listItemToResponse(item))
var tags []string
err := rows.Scan(
&ch.ID, &ch.UserID, &ch.Title, &ch.Type, &ch.AiMode, &ch.Topic,
&ch.Description, &ch.Model, &ch.ProviderConfigID,
&ch.SystemPrompt, &ch.IsArchived, &ch.IsPinned, &ch.Folder, &ch.FolderID, &ch.ProjectID, &ch.WorkspaceID,
scanTags(&tags), scanJSON(&ch.Settings),
&ch.MessageCount, &ch.CreatedAt, &ch.UpdatedAt,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to scan channel"})
return
}
if tags == nil {
tags = []string{}
}
ch.Tags = tags
channels = append(channels, ch)
}
rows.Close()
// v0.23.2: Compute unread counts per channel (safe post-processing)
for i := range channels {
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COUNT(*) FROM messages m
JOIN channel_participants cp ON cp.channel_id = m.channel_id
WHERE cp.channel_id = $1
AND cp.participant_type = 'user' AND cp.participant_id = $2
AND m.created_at > cp.last_read_at
`), channels[i].ID, userID).Scan(&channels[i].UnreadCount)
} }
SafeJSON(c, http.StatusOK, paginatedResponse{ SafeJSON(c, http.StatusOK, paginatedResponse{
@@ -370,6 +258,8 @@ func (h *ChannelHandler) CreateChannel(c *gin.Context) {
} }
} }
ctx := c.Request.Context()
// ── DM dedup: check for existing DM between these two users ── // ── DM dedup: check for existing DM between these two users ──
if channelType == "dm" && len(req.Participants) == 1 { if channelType == "dm" && len(req.Participants) == 1 {
otherUserID := req.Participants[0] otherUserID := req.Participants[0]
@@ -378,162 +268,54 @@ func (h *ChannelHandler) CreateChannel(c *gin.Context) {
return return
} }
// Look for an existing DM channel where both users are participants existingID, _ := h.stores.Channels.FindExistingDM(ctx, userID, otherUserID)
var existingID string
_ = database.DB.QueryRow(database.Q(`
SELECT cp1.channel_id FROM channel_participants cp1
JOIN channel_participants cp2 ON cp1.channel_id = cp2.channel_id
JOIN channels c ON c.id = cp1.channel_id
WHERE c.type = 'dm'
AND cp1.participant_type = 'user' AND cp1.participant_id = $1
AND cp2.participant_type = 'user' AND cp2.participant_id = $2
LIMIT 1
`), userID, otherUserID).Scan(&existingID)
if existingID != "" { if existingID != "" {
// Return existing channel instead of creating duplicate // Return existing channel instead of creating duplicate
var ch channelResponse item, err := h.stores.Channels.GetForUser(ctx, existingID, userID)
var tags []string
err := database.DB.QueryRow(database.Q(`
SELECT id, user_id, title, type, ai_mode, topic,
description, model, provider_config_id,
system_prompt, is_archived, is_pinned, folder, folder_id, project_id, workspace_id,
tags, settings,
created_at, updated_at
FROM channels WHERE id = $1
`), existingID).Scan(
&ch.ID, &ch.UserID, &ch.Title, &ch.Type, &ch.AiMode, &ch.Topic,
&ch.Description, &ch.Model, &ch.ProviderConfigID,
&ch.SystemPrompt, &ch.IsArchived, &ch.IsPinned, &ch.Folder, &ch.FolderID, &ch.ProjectID, &ch.WorkspaceID,
scanTags(&tags), scanJSON(&ch.Settings), &ch.CreatedAt, &ch.UpdatedAt,
)
if err == nil { if err == nil {
if tags == nil { SafeJSON(c, http.StatusOK, listItemToResponse(*item)) // 200, not 201 — existing resource
tags = []string{}
}
ch.Tags = tags
SafeJSON(c, http.StatusOK, ch) // 200, not 201 — existing resource
return return
} }
// If read fails, fall through and create new (shouldn't happen) // If read fails, fall through and create new (shouldn't happen)
} }
} }
// INSERT and retrieve the new row // Build channel model
var ch channelResponse ch := &models.Channel{
var tags []string UserID: userID,
Title: req.Title,
if database.IsSQLite() { Type: channelType,
id := store.NewID() Description: req.Description,
_, err := database.DB.Exec(` Model: req.Model,
INSERT INTO channels (id, user_id, title, type, description, model, SystemPrompt: req.SystemPrompt,
system_prompt, provider_config_id, folder, folder_id, tags, ai_mode) ProviderConfigID: req.ProviderConfigID,
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, FolderID: req.FolderID,
id, userID, req.Title, channelType, req.Description, req.Model,
req.SystemPrompt, req.ProviderConfigID, req.Folder, req.FolderID, writeTagsArg(req.Tags), aiMode,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create channel"})
return
}
// Read back the row
err = database.DB.QueryRow(`
SELECT id, user_id, title, type, ai_mode, topic,
description, model, provider_config_id,
system_prompt, is_archived, is_pinned, folder, folder_id, project_id, workspace_id,
tags, settings,
created_at, updated_at
FROM channels WHERE id = ?`, id).Scan(
&ch.ID, &ch.UserID, &ch.Title, &ch.Type, &ch.AiMode, &ch.Topic,
&ch.Description, &ch.Model, &ch.ProviderConfigID,
&ch.SystemPrompt, &ch.IsArchived, &ch.IsPinned, &ch.Folder, &ch.FolderID, &ch.ProjectID, &ch.WorkspaceID,
scanTags(&tags), scanJSON(&ch.Settings), &ch.CreatedAt, &ch.UpdatedAt,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to read created channel: " + err.Error()})
return
}
} else {
err := database.DB.QueryRow(`
INSERT INTO channels (user_id, title, type, description, model, system_prompt, provider_config_id, folder, folder_id, tags, ai_mode)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)
RETURNING id, user_id, title, type, ai_mode, topic,
description, model, provider_config_id, system_prompt,
is_archived, is_pinned, folder, folder_id, project_id, workspace_id, tags, settings, created_at, updated_at
`, userID, req.Title, channelType, req.Description, req.Model, req.SystemPrompt, req.ProviderConfigID,
req.Folder, req.FolderID, pq.Array(req.Tags), aiMode,
).Scan(
&ch.ID, &ch.UserID, &ch.Title, &ch.Type, &ch.AiMode, &ch.Topic,
&ch.Description, &ch.Model, &ch.ProviderConfigID,
&ch.SystemPrompt, &ch.IsArchived, &ch.IsPinned, &ch.Folder, &ch.FolderID, &ch.ProjectID, &ch.WorkspaceID,
pq.Array(&tags), scanJSON(&ch.Settings), &ch.CreatedAt, &ch.UpdatedAt,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create channel"})
return
}
} }
if tags == nil { dmPartners := []string{}
tags = []string{} if channelType == "dm" {
} dmPartners = req.Participants
ch.Tags = tags
ch.MessageCount = 0
// Auto-create channel_participant for the creator
if database.IsSQLite() {
_, _ = database.DB.Exec(`
INSERT INTO channel_participants (id, channel_id, participant_type, participant_id, role)
VALUES (?, ?, 'user', ?, 'owner')
ON CONFLICT DO NOTHING
`, store.NewID(), ch.ID, userID)
} else {
_, _ = database.DB.Exec(`
INSERT INTO channel_participants (channel_id, participant_type, participant_id, role)
VALUES ($1, 'user', $2, 'owner')
ON CONFLICT DO NOTHING
`, ch.ID, userID)
} }
// v0.23.2: Add DM partner as member participant defaultConfigID := ""
if channelType == "dm" && len(req.Participants) > 0 { if req.ProviderConfigID != nil {
for _, pid := range req.Participants { defaultConfigID = *req.ProviderConfigID
if pid == userID {
continue // skip self (already added as owner)
}
if database.IsSQLite() {
_, _ = database.DB.Exec(`
INSERT INTO channel_participants (id, channel_id, participant_type, participant_id, role)
VALUES (?, ?, 'user', ?, 'member')
ON CONFLICT DO NOTHING
`, store.NewID(), ch.ID, pid)
} else {
_, _ = database.DB.Exec(`
INSERT INTO channel_participants (channel_id, participant_type, participant_id, role)
VALUES ($1, 'user', $2, 'member')
ON CONFLICT DO NOTHING
`, ch.ID, pid)
}
}
} }
// Auto-create channel_model if model specified if err := h.stores.Channels.CreateFull(ctx, ch, req.Folder, req.Tags, aiMode,
if req.Model != "" { userID, dmPartners, req.Model, defaultConfigID); err != nil {
if database.IsSQLite() { c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create channel"})
_, _ = database.DB.Exec(` return
INSERT INTO channel_models (id, channel_id, model_id, provider_config_id, is_default)
VALUES (?, ?, ?, ?, 1)
ON CONFLICT DO NOTHING
`, store.NewID(), ch.ID, req.Model, req.ProviderConfigID)
} else {
_, _ = database.DB.Exec(`
INSERT INTO channel_models (channel_id, model_id, provider_config_id, is_default)
VALUES ($1, $2, $3, true)
ON CONFLICT DO NOTHING
`, ch.ID, req.Model, req.ProviderConfigID)
}
} }
SafeJSON(c, http.StatusCreated, ch) // Fetch full response with message count
item, err := h.stores.Channels.GetForUser(ctx, ch.ID, userID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to read created channel"})
return
}
SafeJSON(c, http.StatusCreated, listItemToResponse(*item))
} }
// ── Get Channel ───────────────────────────── // ── Get Channel ─────────────────────────────
@@ -542,45 +324,13 @@ func (h *ChannelHandler) GetChannel(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
var ch channelResponse item, err := h.stores.Channels.GetForUser(c.Request.Context(), channelID, userID)
var tags []string if err != nil {
err := database.DB.QueryRow(database.Q(`
SELECT c.id, c.user_id, c.title, c.type, c.ai_mode, c.topic,
c.description, c.model, c.provider_config_id,
c.system_prompt, c.is_archived, c.is_pinned, c.folder, c.folder_id, c.project_id, c.workspace_id,
c.tags, c.settings,
COALESCE(mc.cnt, 0) AS message_count,
c.created_at, c.updated_at
FROM channels c
LEFT JOIN (
SELECT channel_id, COUNT(*) AS cnt FROM messages GROUP BY channel_id
) mc ON mc.channel_id = c.id
WHERE c.id = $1 AND (c.user_id = $2 OR c.id IN (
SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = $2
))
`), channelID, userID).Scan(
&ch.ID, &ch.UserID, &ch.Title, &ch.Type, &ch.AiMode, &ch.Topic,
&ch.Description, &ch.Model, &ch.ProviderConfigID,
&ch.SystemPrompt, &ch.IsArchived, &ch.IsPinned, &ch.Folder, &ch.FolderID, &ch.ProjectID, &ch.WorkspaceID,
scanTags(&tags), scanJSON(&ch.Settings),
&ch.MessageCount, &ch.CreatedAt, &ch.UpdatedAt,
)
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return return
} }
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to get channel"})
return
}
if tags == nil { SafeJSON(c, http.StatusOK, listItemToResponse(*item))
tags = []string{}
}
ch.Tags = tags
SafeJSON(c, http.StatusOK, ch)
} }
// ── Update Channel ────────────────────────── // ── Update Channel ──────────────────────────
@@ -588,11 +338,7 @@ func (h *ChannelHandler) GetChannel(c *gin.Context) {
func (h *ChannelHandler) UpdateChannel(c *gin.Context) { func (h *ChannelHandler) UpdateChannel(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
ctx := c.Request.Context()
if database.DB == nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "database unavailable"})
return
}
var req updateChannelRequest var req updateChannelRequest
if err := c.ShouldBindJSON(&req); err != nil { if err := c.ShouldBindJSON(&req); err != nil {
@@ -601,79 +347,62 @@ func (h *ChannelHandler) UpdateChannel(c *gin.Context) {
} }
// Verify ownership // Verify ownership
var ownerID string owns, err := h.stores.Channels.UserOwns(ctx, channelID, userID)
err := database.DB.QueryRow(database.Q(`SELECT user_id FROM channels WHERE id = $1`), channelID).Scan(&ownerID) if err != nil {
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return return
} }
if ownerID != userID { if !owns {
c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"}) c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"})
return return
} }
// Build dynamic UPDATE with ? placeholders (converted for Postgres) // Build fields map for store.Update
setClauses := []string{} fields := map[string]interface{}{}
args := []interface{}{}
addClause := func(col string, val interface{}) {
setClauses = append(setClauses, col+" = ?")
args = append(args, val)
}
if req.Title != nil { if req.Title != nil {
addClause("title", *req.Title) fields["title"] = *req.Title
} }
if req.Description != nil { if req.Description != nil {
addClause("description", *req.Description) fields["description"] = *req.Description
} }
if req.Model != nil { if req.Model != nil {
addClause("model", *req.Model) fields["model"] = *req.Model
} }
if req.SystemPrompt != nil { if req.SystemPrompt != nil {
addClause("system_prompt", *req.SystemPrompt) fields["system_prompt"] = *req.SystemPrompt
} }
if req.ProviderConfigID != nil { if req.ProviderConfigID != nil {
addClause("provider_config_id", *req.ProviderConfigID) if *req.ProviderConfigID == "" {
fields["provider_config_id"] = nil
} else {
fields["provider_config_id"] = *req.ProviderConfigID
}
} }
if req.IsArchived != nil { if req.IsArchived != nil {
addClause("is_archived", *req.IsArchived) fields["is_archived"] = *req.IsArchived
} }
if req.IsPinned != nil { if req.IsPinned != nil {
addClause("is_pinned", *req.IsPinned) fields["is_pinned"] = *req.IsPinned
} }
if req.Folder != nil { if req.Folder != nil {
addClause("folder", *req.Folder) fields["folder"] = *req.Folder
} }
if req.FolderID != nil { if req.FolderID != nil {
if *req.FolderID == "" { if *req.FolderID == "" {
addClause("folder_id", nil) // unbind from folder fields["folder_id"] = nil // unbind from folder
} else { } else {
addClause("folder_id", *req.FolderID) fields["folder_id"] = *req.FolderID
} }
} }
if req.Tags != nil { if req.Tags != nil {
addClause("tags", writeTagsArg(req.Tags)) fields["tags"] = req.Tags
}
if req.Settings != nil {
// Validate settings is well-formed JSON before writing
if !json.Valid([]byte(*req.Settings)) {
c.JSON(http.StatusBadRequest, gin.H{"error": "settings must be valid JSON"})
return
}
// JSONB merge: new settings keys overwrite existing, unmentioned keys preserved
if database.IsSQLite() {
setClauses = append(setClauses, "settings = json_patch(settings, ?)")
} else {
setClauses = append(setClauses, "settings = COALESCE(settings, '{}'::jsonb) || ?::jsonb")
}
args = append(args, []byte(*req.Settings))
} }
if req.WorkspaceID != nil { if req.WorkspaceID != nil {
if *req.WorkspaceID == "" { if *req.WorkspaceID == "" {
addClause("workspace_id", nil) // unbind fields["workspace_id"] = nil // unbind
} else { } else {
addClause("workspace_id", *req.WorkspaceID) fields["workspace_id"] = *req.WorkspaceID
} }
} }
if req.AiMode != nil { if req.AiMode != nil {
@@ -682,29 +411,39 @@ func (h *ChannelHandler) UpdateChannel(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "ai_mode must be auto, mention_only, or off"}) c.JSON(http.StatusBadRequest, gin.H{"error": "ai_mode must be auto, mention_only, or off"})
return return
} }
addClause("ai_mode", mode) fields["ai_mode"] = mode
} }
if req.Topic != nil { if req.Topic != nil {
addClause("topic", *req.Topic) fields["topic"] = *req.Topic
} }
if len(setClauses) == 0 { hasFields := len(fields) > 0
hasSettings := req.Settings != nil
if !hasFields && !hasSettings {
c.JSON(http.StatusBadRequest, gin.H{"error": "no fields to update"}) c.JSON(http.StatusBadRequest, gin.H{"error": "no fields to update"})
return return
} }
query := "UPDATE channels SET " + strings.Join(setClauses, ", ") if hasSettings {
query += " WHERE id = ? AND user_id = ?" if !json.Valid([]byte(*req.Settings)) {
args = append(args, channelID, userID) c.JSON(http.StatusBadRequest, gin.H{"error": "settings must be valid JSON"})
return
if !database.IsSQLite() { }
query = convertPlaceholders(query)
} }
_, err = database.DB.Exec(query, args...) if hasFields {
if err != nil { if err := h.stores.Channels.Update(ctx, channelID, fields); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update channel"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update channel"})
return return
}
}
if hasSettings {
if err := h.stores.Channels.MergeSettings(ctx, channelID, json.RawMessage(*req.Settings)); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update settings"})
return
}
} }
// Return updated channel // Return updated channel
@@ -717,17 +456,12 @@ func (h *ChannelHandler) DeleteChannel(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
result, err := database.DB.Exec( n, err := h.stores.Channels.DeleteByOwner(c.Request.Context(), channelID, userID)
database.Q(`DELETE FROM channels WHERE id = $1 AND user_id = $2`),
channelID, userID,
)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete channel"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete channel"})
return return
} }
if n == 0 {
rows, _ := result.RowsAffected()
if rows == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return return
} }
@@ -746,30 +480,9 @@ func (h *ChannelHandler) MarkRead(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
// Update the read cursor timestamp (last_read_at exists since migration 005) err := h.stores.Channels.MarkRead(c.Request.Context(), channelID, userID)
_, err := database.DB.ExecContext(c.Request.Context(), database.Q(`
UPDATE channel_participants
SET last_read_at = NOW()
WHERE channel_id = $1 AND participant_type = 'user' AND participant_id = $2
`), channelID, userID)
if err != nil { if err != nil {
// Participant row may not exist for legacy direct chats — that's OK // Participant row may not exist for legacy direct chats — that's OK
c.JSON(http.StatusOK, gin.H{"ok": true})
return
}
// Best-effort: also set last_read_message_id if the column exists (migration 017)
var latestMsgID *string
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT id FROM messages WHERE channel_id = $1 ORDER BY created_at DESC LIMIT 1
`), channelID).Scan(&latestMsgID)
if latestMsgID != nil {
// This may fail if 017 hasn't been applied — that's fine, last_read_at is primary
_, _ = database.DB.ExecContext(c.Request.Context(), database.Q(`
UPDATE channel_participants
SET last_read_message_id = $1
WHERE channel_id = $2 AND participant_type = 'user' AND participant_id = $3
`), *latestMsgID, channelID, userID)
} }
c.JSON(http.StatusOK, gin.H{"ok": true}) c.JSON(http.StatusOK, gin.H{"ok": true})

View File

@@ -18,6 +18,7 @@ import (
"git.gobha.me/xcaliber/chat-switchboard/crypto" "git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/events" "git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/filters"
"git.gobha.me/xcaliber/chat-switchboard/health" "git.gobha.me/xcaliber/chat-switchboard/health"
"git.gobha.me/xcaliber/chat-switchboard/knowledge" "git.gobha.me/xcaliber/chat-switchboard/knowledge"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
@@ -55,6 +56,7 @@ type CompletionHandler struct {
health HealthRecorder // provider health tracking (v0.22.0, nil = disabled) health HealthRecorder // provider health tracking (v0.22.0, nil = disabled)
healthStore HealthStatusQuerier // health status queries for routing (v0.22.2, 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) router *routing.Evaluator // routing policy evaluator (v0.22.2, nil = disabled)
filterChain *filters.Chain // pre-completion filter chain (v0.29.0, nil = disabled)
} }
// HealthRecorder is the interface for recording provider call outcomes. // HealthRecorder is the interface for recording provider call outcomes.
@@ -94,6 +96,11 @@ func (h *CompletionHandler) SetRoutingEvaluator(r *routing.Evaluator) {
h.router = r h.router = r
} }
// SetFilterChain attaches the pre-completion filter chain (v0.29.0).
func (h *CompletionHandler) SetFilterChain(fc *filters.Chain) {
h.filterChain = fc
}
// evaluateRouting applies routing policies to select the best provider config // evaluateRouting applies routing policies to select the best provider config
// for this request. Returns the winning config details and a routing decision // for this request. Returns the winning config details and a routing decision
// for observability. If routing is disabled or no policies match, returns // for observability. If routing is disabled or no policies match, returns
@@ -430,7 +437,7 @@ func (h *CompletionHandler) Complete(c *gin.Context) {
} }
// ── Team policy: require_private_providers ── // ── Team policy: require_private_providers ──
if err := enforcePrivateProviderPolicy(userID, configID); err != nil { if err := enforcePrivateProviderPolicy(c.Request.Context(), h.stores, userID, configID); err != nil {
c.JSON(http.StatusForbidden, gin.H{"error": err.Error()}) c.JSON(http.StatusForbidden, gin.H{"error": err.Error()})
return return
} }
@@ -1092,7 +1099,7 @@ func escapeJSON(s string) string {
// getModelCapabilities looks up capabilities from model_catalog DB, // getModelCapabilities looks up capabilities from model_catalog DB,
// then overlays with known model defaults and heuristic detection. // then overlays with known model defaults and heuristic detection.
func (h *CompletionHandler) getModelCapabilities(c *gin.Context, model, apiConfigID string) models.ModelCapabilities { func (h *CompletionHandler) getModelCapabilities(c *gin.Context, model, apiConfigID string) models.ModelCapabilities {
return ResolveModelCaps(c, model, apiConfigID) return ResolveModelCaps(h.stores.Catalog, model, apiConfigID)
} }
// ── Multimodal Assembly ───────────────────── // ── Multimodal Assembly ─────────────────────
@@ -1513,7 +1520,7 @@ func extractFirstMention(content string) string {
// Priority: request.provider_config_id → chat.provider_config_id → user's first active config // 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) { func (h *CompletionHandler) resolveConfig(userID string, channelID string, req completionRequest) (providers.ProviderConfig, string, string, string, string, error) {
res, err := ResolveProviderConfig(h.vault, userID, channelID, req.ProviderConfigID, req.Model) res, err := ResolveProviderConfig(h.stores, h.vault, userID, channelID, req.ProviderConfigID, req.Model)
if err != nil { if err != nil {
return providers.ProviderConfig{}, "", "", "", "", err return providers.ProviderConfig{}, "", "", "", "", err
} }
@@ -1576,12 +1583,22 @@ func (h *CompletionHandler) loadConversation(channelID, userID, personaSystemPro
} }
} }
// ── KB hint (nudge LLM to use kb_search when KBs are active) ── // ── Pre-completion filter chain (v0.29.0) ──
if kbHint := BuildKBHint(context.Background(), h.stores, channelID, userID, personaID); kbHint != "" { // Runs registered filters (KB auto-inject, future extension filters).
messages = append(messages, providers.Message{ // Replaces the inline BuildKBHint call. Each filter contributes zero
Role: "system", // or more system messages. Failures are logged and skipped.
Content: kbHint, 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) ── // ── Memory hint (inject known facts about this user — v0.18.0) ──

View File

@@ -0,0 +1,213 @@
package handlers
import (
"net/http"
"github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ExtPermHandler serves extension permission management endpoints.
// v0.29.0 CS2: Admin reviews and grants/revokes declared permissions.
type ExtPermHandler struct {
stores store.Stores
}
func NewExtPermHandler(stores store.Stores) *ExtPermHandler {
return &ExtPermHandler{stores: stores}
}
// ListPackagePermissions returns all declared permissions for a package.
// GET /api/v1/admin/extensions/:id/permissions
func (h *ExtPermHandler) ListPackagePermissions(c *gin.Context) {
pkgID := c.Param("id")
perms, err := h.stores.ExtPermissions.ListForPackage(c.Request.Context(), pkgID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list permissions"})
return
}
c.JSON(http.StatusOK, gin.H{"data": perms})
}
// GrantPermission grants a single declared permission for a package.
// POST /api/v1/admin/extensions/:id/permissions/:perm/grant
func (h *ExtPermHandler) GrantPermission(c *gin.Context) {
pkgID := c.Param("id")
perm := c.Param("perm")
userID := c.GetString("user_id")
if !models.ValidExtensionPermissions[perm] {
c.JSON(http.StatusBadRequest, gin.H{"error": "unrecognized permission: " + perm})
return
}
if err := h.stores.ExtPermissions.Grant(c.Request.Context(), pkgID, perm, userID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to grant permission"})
return
}
// Check if all permissions are now granted → activate package
h.maybeActivate(c, pkgID)
c.JSON(http.StatusOK, gin.H{"status": "granted", "package_id": pkgID, "permission": perm})
}
// RevokePermission revokes a single permission for a package.
// POST /api/v1/admin/extensions/:id/permissions/:perm/revoke
func (h *ExtPermHandler) RevokePermission(c *gin.Context) {
pkgID := c.Param("id")
perm := c.Param("perm")
if err := h.stores.ExtPermissions.Revoke(c.Request.Context(), pkgID, perm); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to revoke permission"})
return
}
// Suspend if a previously-active package loses a grant
h.maybeSuspend(c, pkgID)
c.JSON(http.StatusOK, gin.H{"status": "revoked", "package_id": pkgID, "permission": perm})
}
// GrantAllPermissions grants all declared permissions for a package.
// POST /api/v1/admin/extensions/:id/permissions/grant-all
func (h *ExtPermHandler) GrantAllPermissions(c *gin.Context) {
pkgID := c.Param("id")
userID := c.GetString("user_id")
if err := h.stores.ExtPermissions.GrantAll(c.Request.Context(), pkgID, userID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to grant all permissions"})
return
}
// Activate the package now that all permissions are granted
h.maybeActivate(c, pkgID)
c.JSON(http.StatusOK, gin.H{"status": "granted_all", "package_id": pkgID})
}
// ReviewPackage returns the package with its declared permissions.
// GET /api/v1/admin/extensions/:id/review
func (h *ExtPermHandler) ReviewPackage(c *gin.Context) {
pkgID := c.Param("id")
pkg, err := h.stores.Packages.Get(c.Request.Context(), pkgID)
if err != nil || pkg == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "package not found"})
return
}
perms, err := h.stores.ExtPermissions.ListForPackage(c.Request.Context(), pkgID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list permissions"})
return
}
c.JSON(http.StatusOK, gin.H{
"package": pkg,
"permissions": perms,
})
}
// ── Lifecycle helpers ────────────────────────
// maybeActivate transitions a pending_review package to active
// if all declared permissions have been granted.
func (h *ExtPermHandler) maybeActivate(c *gin.Context, pkgID string) {
pkg, err := h.stores.Packages.Get(c.Request.Context(), pkgID)
if err != nil || pkg == nil {
return
}
if pkg.Status != models.PackageStatusPendingReview {
return
}
perms, err := h.stores.ExtPermissions.ListForPackage(c.Request.Context(), pkgID)
if err != nil {
return
}
allGranted := true
for _, p := range perms {
if !p.Granted {
allGranted = false
break
}
}
if allGranted {
_ = h.stores.Packages.SetStatus(c.Request.Context(), pkgID, models.PackageStatusActive)
}
}
// maybeSuspend transitions an active package to suspended
// if a required permission was revoked.
func (h *ExtPermHandler) maybeSuspend(c *gin.Context, pkgID string) {
pkg, err := h.stores.Packages.Get(c.Request.Context(), pkgID)
if err != nil || pkg == nil {
return
}
if pkg.Status != models.PackageStatusActive {
return
}
perms, err := h.stores.ExtPermissions.ListForPackage(c.Request.Context(), pkgID)
if err != nil || len(perms) == 0 {
return
}
for _, p := range perms {
if !p.Granted {
_ = h.stores.Packages.SetStatus(c.Request.Context(), pkgID, models.PackageStatusSuspended)
return
}
}
}
// ── Manifest Permission Parsing ──────────────
// SyncManifestPermissions extracts the "permissions" array from a
// package manifest and declares them in the store. Returns the list
// of declared permissions. If the manifest declares permissions,
// the package status is set to pending_review.
//
// Called by AdminInstallExtension and InstallPackage handlers.
func SyncManifestPermissions(c *gin.Context, stores store.Stores, pkgID string, manifest map[string]any) []string {
if stores.ExtPermissions == nil {
return nil
}
raw, ok := manifest["permissions"]
if !ok {
return nil
}
arr, ok := raw.([]interface{})
if !ok {
return nil
}
var perms []string
for _, v := range arr {
if s, ok := v.(string); ok && models.ValidExtensionPermissions[s] {
perms = append(perms, s)
}
}
if len(perms) == 0 {
return nil
}
if err := stores.ExtPermissions.DeclareForPackage(c.Request.Context(), pkgID, perms); err != nil {
return nil
}
// Package needs review before activation
_ = stores.Packages.SetStatus(c.Request.Context(), pkgID, models.PackageStatusPendingReview)
return perms
}

View File

@@ -0,0 +1,105 @@
package handlers
import (
"net/http"
"github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ExtSecretsHandler serves extension secret management endpoints.
// v0.29.0 CS3: Admin sets key-value secrets that Starlark extensions
// can read via the secrets.get() module.
//
// Secrets are stored in GlobalConfig under key "ext_secrets:{packageID}"
// as a JSON map: {"api_key": "sk-...", "webhook_token": "tok-..."}.
type ExtSecretsHandler struct {
stores store.Stores
}
func NewExtSecretsHandler(stores store.Stores) *ExtSecretsHandler {
return &ExtSecretsHandler{stores: stores}
}
// GetSecrets returns the secret keys (not values) for a package.
// GET /api/v1/admin/extensions/:id/secrets
func (h *ExtSecretsHandler) GetSecrets(c *gin.Context) {
pkgID := c.Param("id")
configKey := "ext_secrets:" + pkgID
configVal, err := h.stores.GlobalConfig.Get(c.Request.Context(), configKey)
if err != nil {
// Not found = empty secrets, not an error
c.JSON(http.StatusOK, gin.H{"data": map[string]string{}})
return
}
// Return only keys, not values (security: don't expose secrets in GET)
keys := make([]string, 0, len(configVal))
for k := range configVal {
keys = append(keys, k)
}
c.JSON(http.StatusOK, gin.H{"data": gin.H{
"package_id": pkgID,
"keys": keys,
}})
}
// SetSecrets upserts secrets for a package.
// PUT /api/v1/admin/extensions/:id/secrets
// Body: {"secrets": {"api_key": "sk-...", "webhook_token": "tok-..."}}
func (h *ExtSecretsHandler) SetSecrets(c *gin.Context) {
pkgID := c.Param("id")
userID := c.GetString("user_id")
// Verify package exists
pkg, err := h.stores.Packages.Get(c.Request.Context(), pkgID)
if err != nil || pkg == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "package not found"})
return
}
var body struct {
Secrets map[string]string `json:"secrets" binding:"required"`
}
if err := c.ShouldBindJSON(&body); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid request: " + err.Error()})
return
}
// Store as GlobalConfig JSONMap
configKey := "ext_secrets:" + pkgID
configVal := models.JSONMap{}
for k, v := range body.Secrets {
configVal[k] = v
}
if err := h.stores.GlobalConfig.Set(c.Request.Context(), configKey, configVal, userID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to save secrets"})
return
}
c.JSON(http.StatusOK, gin.H{
"status": "saved",
"package_id": pkgID,
"key_count": len(body.Secrets),
})
}
// DeleteSecrets removes all secrets for a package.
// DELETE /api/v1/admin/extensions/:id/secrets
func (h *ExtSecretsHandler) DeleteSecrets(c *gin.Context) {
pkgID := c.Param("id")
userID := c.GetString("user_id")
configKey := "ext_secrets:" + pkgID
if err := h.stores.GlobalConfig.Set(c.Request.Context(), configKey, models.JSONMap{}, userID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete secrets"})
return
}
c.JSON(http.StatusOK, gin.H{"status": "deleted", "package_id": pkgID})
}

View File

@@ -206,6 +206,10 @@ func (h *ExtensionHandler) AdminInstallExtension(c *gin.Context) {
return return
} }
// v0.29.0: Parse manifest permissions and declare them.
// If permissions are declared, package moves to pending_review.
SyncManifestPermissions(c, h.stores, pkg.ID, manifestMap)
c.JSON(201, gin.H{"data": pkg}) c.JSON(201, gin.H{"data": pkg})
} }

View File

@@ -12,7 +12,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/extraction" "git.gobha.me/xcaliber/chat-switchboard/extraction"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/storage" "git.gobha.me/xcaliber/chat-switchboard/storage"
@@ -162,10 +161,8 @@ func (h *FileHandler) Upload(c *gin.Context) {
return return
} }
// Update storage_key in PG // Update storage_key
database.DB.ExecContext(c.Request.Context(), h.stores.Files.UpdateStorageKey(c.Request.Context(), att.ID, att.StorageKey)
database.Q(`UPDATE files SET storage_key = $1 WHERE id = $2`),
att.StorageKey, att.ID)
// For images, mark extraction as not needed (complete immediately) // For images, mark extraction as not needed (complete immediately)
if strings.HasPrefix(contentType, "image/") { if strings.HasPrefix(contentType, "image/") {
@@ -463,14 +460,12 @@ func (h *FileHandler) CleanupChannelStorage(channelID string) {
// verifyChannelAccess checks that the requesting user owns the channel. // verifyChannelAccess checks that the requesting user owns the channel.
// Future RBAC (v0.20.0): replace with rbac.Can(userID, channelID, permission). // Future RBAC (v0.20.0): replace with rbac.Can(userID, channelID, permission).
func (h *FileHandler) verifyChannelAccess(c *gin.Context, channelID, userID string) bool { func (h *FileHandler) verifyChannelAccess(c *gin.Context, channelID, userID string) bool {
var ownerID string owns, err := h.stores.Channels.UserOwns(c.Request.Context(), channelID, userID)
err := database.DB.QueryRowContext(c.Request.Context(),
database.Q(`SELECT user_id FROM channels WHERE id = $1`), channelID).Scan(&ownerID)
if err != nil { if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return false return false
} }
if ownerID != userID { if !owns {
// Check if user is admin (admins can access any channel) // Check if user is admin (admins can access any channel)
role, _ := c.Get("role") role, _ := c.Get("role")
if role != "admin" { if role != "admin" {
@@ -516,16 +511,9 @@ func (h *FileHandler) UploadToProject(c *gin.Context) {
// Verify project access (user owns or is team member; admins bypass) // Verify project access (user owns or is team member; admins bypass)
if c.GetString("role") != "admin" { if c.GetString("role") != "admin" {
var exists bool teamIDs, _ := h.stores.Teams.GetUserTeamIDs(c.Request.Context(), userID)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(` ok, err := h.stores.Projects.UserCanAccess(c.Request.Context(), userID, projectID, teamIDs)
SELECT EXISTS( if err != nil || !ok {
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"}) c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return return
} }
@@ -597,16 +585,9 @@ func (h *FileHandler) ListByProject(c *gin.Context) {
// Admins bypass project ownership/membership checks // Admins bypass project ownership/membership checks
if c.GetString("role") != "admin" { if c.GetString("role") != "admin" {
var exists bool teamIDs, _ := h.stores.Teams.GetUserTeamIDs(c.Request.Context(), userID)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(` ok, err := h.stores.Projects.UserCanAccess(c.Request.Context(), userID, projectID, teamIDs)
SELECT EXISTS( if err != nil || !ok {
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"}) c.JSON(http.StatusForbidden, gin.H{"error": "project not found or access denied"})
return return
} }

View File

@@ -2,61 +2,34 @@ package handlers
// folders.go — Chat folder CRUD (v0.23.1) // folders.go — Chat folder CRUD (v0.23.1)
// //
// Folders are user-scoped groupings for personal (direct) chats. // v0.29.0: Raw SQL replaced with FolderStore methods.
// Routes:
// GET /api/v1/folders
// POST /api/v1/folders
// PUT /api/v1/folders/:id
// DELETE /api/v1/folders/:id
import ( import (
"net/http" "net/http"
"strings" "strings"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
type FolderHandler struct{} type FolderHandler struct {
stores store.Stores
}
func NewFolderHandler() *FolderHandler { return &FolderHandler{} } func NewFolderHandler(s store.Stores) *FolderHandler {
return &FolderHandler{stores: s}
type folderRow struct {
ID string `json:"id"`
Name string `json:"name"`
ParentID *string `json:"parent_id,omitempty"`
SortOrder int `json:"sort_order"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
} }
func (h *FolderHandler) List(c *gin.Context) { func (h *FolderHandler) List(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(`
SELECT id, name, parent_id, sort_order, created_at, updated_at folders, err := h.stores.Folders.List(c.Request.Context(), userID)
FROM folders
WHERE user_id = $1
ORDER BY sort_order, name
`), userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list folders"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list folders"})
return return
} }
defer rows.Close()
var folders []folderRow
for rows.Next() {
var f folderRow
if err := rows.Scan(&f.ID, &f.Name, &f.ParentID, &f.SortOrder, &f.CreatedAt, &f.UpdatedAt); err != nil {
continue
}
folders = append(folders, f)
}
if folders == nil {
folders = []folderRow{}
}
c.JSON(http.StatusOK, gin.H{"folders": folders}) c.JSON(http.StatusOK, gin.H{"folders": folders})
} }
@@ -76,14 +49,12 @@ func (h *FolderHandler) Create(c *gin.Context) {
return return
} }
var f folderRow f := &models.Folder{
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(` UserID: userID,
INSERT INTO folders (user_id, name, sort_order) Name: req.Name,
VALUES ($1, $2, $3) SortOrder: req.SortOrder,
RETURNING id, name, parent_id, sort_order, created_at, updated_at }
`), userID, req.Name, req.SortOrder).Scan( if err := h.stores.Folders.Create(c.Request.Context(), f); err != nil {
&f.ID, &f.Name, &f.ParentID, &f.SortOrder, &f.CreatedAt, &f.UpdatedAt)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create folder"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create folder"})
return return
} }
@@ -105,17 +76,12 @@ func (h *FolderHandler) Update(c *gin.Context) {
req.Name = strings.TrimSpace(req.Name) req.Name = strings.TrimSpace(req.Name)
} }
res, err := database.DB.ExecContext(c.Request.Context(), database.Q(` n, err := h.stores.Folders.Update(c.Request.Context(), folderID, userID, req.Name, req.SortOrder)
UPDATE folders
SET name = COALESCE(NULLIF($3, ''), name),
sort_order = COALESCE($4, sort_order)
WHERE id = $1 AND user_id = $2
`), folderID, userID, req.Name, req.SortOrder)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update folder"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update folder"})
return return
} }
if n, _ := res.RowsAffected(); n == 0 { if n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "folder not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "folder not found"})
return return
} }
@@ -127,18 +93,14 @@ func (h *FolderHandler) Delete(c *gin.Context) {
folderID := c.Param("id") folderID := c.Param("id")
// Unassign chats before deleting folder // Unassign chats before deleting folder
_, _ = database.DB.ExecContext(c.Request.Context(), database.Q(` _ = h.stores.Folders.UnassignChannels(c.Request.Context(), folderID, userID)
UPDATE channels SET folder_id = NULL WHERE folder_id = $1 AND user_id = $2
`), folderID, userID)
res, err := database.DB.ExecContext(c.Request.Context(), database.Q(` n, err := h.stores.Folders.Delete(c.Request.Context(), folderID, userID)
DELETE FROM folders WHERE id = $1 AND user_id = $2
`), folderID, userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete folder"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete folder"})
return return
} }
if n, _ := res.RowsAffected(); n == 0 { if n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "folder not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "folder not found"})
return return
} }

View File

@@ -24,6 +24,7 @@ import (
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
postgres "git.gobha.me/xcaliber/chat-switchboard/store/postgres" postgres "git.gobha.me/xcaliber/chat-switchboard/store/postgres"
sqlite "git.gobha.me/xcaliber/chat-switchboard/store/sqlite" sqlite "git.gobha.me/xcaliber/chat-switchboard/store/sqlite"
"git.gobha.me/xcaliber/chat-switchboard/treepath"
) )
// ── Test Harness ──────────────────────────── // ── Test Harness ────────────────────────────
@@ -128,6 +129,7 @@ func setupHarness(t *testing.T) *testHarness {
} else { } else {
stores = postgres.NewStores(database.TestDB) stores = postgres.NewStores(database.TestDB)
} }
treepath.Stores = &stores
userCache := middleware.NewUserStatusCache() userCache := middleware.NewUserStatusCache()
// Roles resolver (nil vault — test-fire won't work, but CRUD will) // Roles resolver (nil vault — test-fire won't work, but CRUD will)
@@ -175,7 +177,7 @@ func setupHarness(t *testing.T) *testHarness {
protected.GET("/teams/mine", teams.MyTeams) protected.GET("/teams/mine", teams.MyTeams)
teamScoped := protected.Group("/teams/:teamId") teamScoped := protected.Group("/teams/:teamId")
teamScoped.Use(middleware.RequireTeamAdmin()) teamScoped.Use(middleware.RequireTeamAdmin(stores.Teams))
{ {
// Team members (team admin self-service) // Team members (team admin self-service)
teamScoped.GET("/members", teams.ListMembers) teamScoped.GET("/members", teams.ListMembers)
@@ -235,7 +237,7 @@ func setupHarness(t *testing.T) *testHarness {
// users bypass the DB query so these tests work on both dialects when // users bypass the DB query so these tests work on both dialects when
// the caller is an admin. Non-admin member tests require Postgres. // the caller is an admin. Non-admin member tests require Postgres.
teamMemberRoutes := protected.Group("/teams/:teamId") teamMemberRoutes := protected.Group("/teams/:teamId")
teamMemberRoutes.Use(middleware.RequireTeamMember()) teamMemberRoutes.Use(middleware.RequireTeamMember(stores.Teams))
{ {
teamMemberTaskH := NewTaskHandler(stores) teamMemberTaskH := NewTaskHandler(stores)
teamMemberRoutes.GET("/tasks", teamMemberTaskH.ListTeamTasks) teamMemberRoutes.GET("/tasks", teamMemberTaskH.ListTeamTasks)
@@ -247,7 +249,7 @@ func setupHarness(t *testing.T) *testHarness {
protected.GET("/usage", usage.PersonalUsage) protected.GET("/usage", usage.PersonalUsage)
// Profile / Settings // Profile / Settings
settings := NewSettingsHandler(nil) settings := NewSettingsHandler(stores, nil)
protected.GET("/profile", settings.GetProfile) protected.GET("/profile", settings.GetProfile)
protected.PUT("/profile", settings.UpdateProfile) protected.PUT("/profile", settings.UpdateProfile)
protected.POST("/profile/password", settings.ChangePassword) protected.POST("/profile/password", settings.ChangePassword)
@@ -270,7 +272,7 @@ func setupHarness(t *testing.T) *testHarness {
protected.PUT("/personas/:id/tool-grants", personas.SetPersonaToolGrants) // v0.25.0 protected.PUT("/personas/:id/tool-grants", personas.SetPersonaToolGrants) // v0.25.0
// Persona Groups // Persona Groups
pgH := NewPersonaGroupHandler() pgH := NewPersonaGroupHandler(stores)
protected.GET("/persona-groups", pgH.List) protected.GET("/persona-groups", pgH.List)
protected.POST("/persona-groups", pgH.Create) protected.POST("/persona-groups", pgH.Create)
protected.GET("/persona-groups/:id", pgH.Get) protected.GET("/persona-groups/:id", pgH.Get)
@@ -288,7 +290,7 @@ func setupHarness(t *testing.T) *testHarness {
protected.DELETE("/notes/:id", notes.Delete) protected.DELETE("/notes/:id", notes.Delete)
// Channels // Channels
channels := NewChannelHandler() channels := NewChannelHandler(stores)
protected.GET("/channels", channels.ListChannels) protected.GET("/channels", channels.ListChannels)
protected.POST("/channels", channels.CreateChannel) protected.POST("/channels", channels.CreateChannel)
protected.GET("/channels/:id", channels.GetChannel) protected.GET("/channels/:id", channels.GetChannel)
@@ -386,8 +388,8 @@ func setupHarness(t *testing.T) *testHarness {
admin.POST("/personas", personas.CreateAdminPersona) admin.POST("/personas", personas.CreateAdminPersona)
admin.PUT("/personas/:id", personas.UpdateAdminPersona) admin.PUT("/personas/:id", personas.UpdateAdminPersona)
admin.DELETE("/personas/:id", personas.DeleteAdminPersona) admin.DELETE("/personas/:id", personas.DeleteAdminPersona)
admin.POST("/personas/:id/avatar", UploadPersonaAvatar) admin.POST("/personas/:id/avatar", func(c *gin.Context) { UploadPersonaAvatar(stores.Personas, c) })
admin.DELETE("/personas/:id/avatar", DeletePersonaAvatar) admin.DELETE("/personas/:id/avatar", func(c *gin.Context) { DeletePersonaAvatar(stores.Personas, c) })
admin.GET("/personas/:id/knowledge-bases", personas.GetPersonaKBs) // v0.17.0 admin.GET("/personas/:id/knowledge-bases", personas.GetPersonaKBs) // v0.17.0
admin.PUT("/personas/:id/knowledge-bases", personas.SetPersonaKBs) // v0.17.0 admin.PUT("/personas/:id/knowledge-bases", personas.SetPersonaKBs) // v0.17.0
admin.GET("/personas/:id/tool-grants", personas.GetPersonaToolGrants) // v0.25.0 admin.GET("/personas/:id/tool-grants", personas.GetPersonaToolGrants) // v0.25.0

View File

@@ -847,6 +847,10 @@ func (h *KnowledgeBaseHandler) userCanAccess(kb *models.KnowledgeBase, userID st
// BuildKBHint returns a system prompt fragment listing active KBs for a // BuildKBHint returns a system prompt fragment listing active KBs for a
// channel, or empty string if none are active. Called by the completion // channel, or empty string if none are active. Called by the completion
// handler to nudge the LLM to use kb_search. // handler to nudge the LLM to use kb_search.
//
// Deprecated: v0.29.0 — replaced by filters.KBInjectFilter in the
// pre-completion filter chain. Kept for backward compatibility with
// any callers outside the main completion path. Will be removed in v0.30.0.
func BuildKBHint(ctx context.Context, stores store.Stores, channelID, userID, personaID string) string { func BuildKBHint(ctx context.Context, stores store.Stores, channelID, userID, personaID string) string {
teamIDs, _ := stores.Teams.GetUserTeamIDs(ctx, userID) teamIDs, _ := stores.Teams.GetUserTeamIDs(ctx, userID)

View File

@@ -12,12 +12,13 @@ import (
capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities" capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities"
"git.gobha.me/xcaliber/chat-switchboard/crypto" "git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/events" "git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/providers" "git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/storage" "git.gobha.me/xcaliber/chat-switchboard/storage"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
"git.gobha.me/xcaliber/chat-switchboard/tools" "git.gobha.me/xcaliber/chat-switchboard/tools"
"git.gobha.me/xcaliber/chat-switchboard/treepath"
) )
// ── Request / Response types ──────────────── // ── Request / Response types ────────────────
@@ -50,12 +51,12 @@ type editRequest struct {
} }
type regenerateRequest struct { type regenerateRequest struct {
Model string `json:"model,omitempty"` Model string `json:"model,omitempty"`
PersonaID string `json:"persona_id,omitempty"` PersonaID string `json:"persona_id,omitempty"`
ProviderConfigID string `json:"provider_config_id,omitempty"` ProviderConfigID string `json:"provider_config_id,omitempty"`
MaxTokens int `json:"max_tokens,omitempty"` MaxTokens int `json:"max_tokens,omitempty"`
Temperature *float64 `json:"temperature,omitempty"` Temperature *float64 `json:"temperature,omitempty"`
DisabledTools []string `json:"disabled_tools,omitempty"` DisabledTools []string `json:"disabled_tools,omitempty"`
} }
type cursorRequest struct { type cursorRequest struct {
@@ -77,9 +78,6 @@ func NewMessageHandler(vault *crypto.KeyResolver, stores store.Stores, hub *even
// ── List Messages (flat, all branches) ────── // ── List Messages (flat, all branches) ──────
// GET /channels/:id/messages // GET /channels/:id/messages
//
// Returns all live messages across all branches (useful for admin/debug).
// Frontend should prefer /channels/:id/path for rendering.
func (h *MessageHandler) ListMessages(c *gin.Context) { func (h *MessageHandler) ListMessages(c *gin.Context) {
page, perPage, offset := parsePagination(c) page, perPage, offset := parsePagination(c)
@@ -93,62 +91,37 @@ func (h *MessageHandler) ListMessages(c *gin.Context) {
c.JSON(http.StatusForbidden, gin.H{"error": "session not authorized for this channel"}) c.JSON(http.StatusForbidden, gin.H{"error": "session not authorized for this channel"})
return return
} }
} else if !userOwnsChannel(c, channelID, userID) { } else if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
var total int ctx := c.Request.Context()
err := database.DB.QueryRow( msgs, total, err := h.stores.Messages.ListWithSenderInfo(ctx, channelID, perPage, offset)
database.Q(`SELECT COUNT(*) FROM messages WHERE channel_id = $1 AND deleted_at IS NULL`),
channelID,
).Scan(&total)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to count messages"})
return
}
rows, err := database.DB.Query(database.Q(`
SELECT m.id, m.channel_id, m.role, m.content, m.model, m.tokens_used, m.parent_id,
m.sibling_index, m.participant_type, m.participant_id,
CASE WHEN m.participant_type = 'user' THEN COALESCE(u.display_name, u.username)
WHEN m.participant_type = 'persona' THEN p.name
ELSE NULL END AS sender_name,
CASE WHEN m.participant_type = 'user' THEN u.avatar_url
WHEN m.participant_type = 'persona' THEN p.avatar
ELSE NULL END AS sender_avatar,
m.created_at
FROM messages m
LEFT JOIN users u ON m.participant_type = 'user' AND m.participant_id = u.id::text
LEFT JOIN personas p ON m.participant_type = 'persona' AND m.participant_id = p.id::text
WHERE m.channel_id = $1 AND m.deleted_at IS NULL
ORDER BY m.created_at ASC
LIMIT $2 OFFSET $3
`), channelID, perPage, offset)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list messages"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list messages"})
return return
} }
defer rows.Close()
messages := make([]messageResponse, 0) messages := make([]messageResponse, 0, len(msgs))
for rows.Next() { for _, m := range msgs {
var msg messageResponse messages = append(messages, messageResponse{
err := rows.Scan( ID: m.ID,
&msg.ID, &msg.ChannelID, &msg.Role, &msg.Content, ChannelID: m.ChannelID,
&msg.Model, &msg.TokensUsed, &msg.ParentID, Role: m.Role,
&msg.SiblingIndex, &msg.ParticipantType, &msg.ParticipantID, Content: m.Content,
&msg.SenderName, &msg.SenderAvatar, Model: m.Model,
&msg.CreatedAt, TokensUsed: m.TokensUsed,
) ParentID: m.ParentID,
if err != nil { SiblingIndex: m.SiblingIndex,
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to scan message"}) ParticipantType: m.ParticipantType,
return ParticipantID: m.ParticipantID,
} SenderName: m.SenderName,
messages = append(messages, msg) SenderAvatar: m.SenderAvatar,
CreatedAt: m.CreatedAt,
})
} }
rows.Close() // release connection before sibling queries
// Enrich with sibling counts (requires DB access, must happen after rows are closed) // Enrich with sibling counts
for i := range messages { for i := range messages {
messages[i].SiblingCount = getSiblingCount(channelID, messages[i].ParentID) messages[i].SiblingCount = getSiblingCount(channelID, messages[i].ParentID)
} }
@@ -164,15 +137,12 @@ func (h *MessageHandler) ListMessages(c *gin.Context) {
// ── Active Path ───────────────────────────── // ── Active Path ─────────────────────────────
// GET /channels/:id/path // GET /channels/:id/path
//
// Returns the active branch from root → leaf with sibling metadata
// at each node. This is what the frontend renders.
func (h *MessageHandler) GetActivePath(c *gin.Context) { func (h *MessageHandler) GetActivePath(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
if !userOwnsChannel(c, channelID, userID) { if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
@@ -205,12 +175,11 @@ func (h *MessageHandler) CreateMessage(c *gin.Context) {
c.JSON(http.StatusForbidden, gin.H{"error": "session not authorized for this channel"}) c.JSON(http.StatusForbidden, gin.H{"error": "session not authorized for this channel"})
return return
} }
} else if !userOwnsChannel(c, channelID, userID) { } else if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
// Use cursor for parent, compute sibling_index // Use cursor for parent, compute sibling_index
// Sessions use the same cursor logic (empty userID = channel-level latest)
effectiveUserID := userID effectiveUserID := userID
if isSessionAuth(c) { if isSessionAuth(c) {
effectiveUserID = c.GetString("session_id") effectiveUserID = c.GetString("session_id")
@@ -231,63 +200,43 @@ func (h *MessageHandler) CreateMessage(c *gin.Context) {
} }
} }
var msg messageResponse msg := &models.Message{
if database.IsSQLite() { ChannelID: channelID,
newID := store.NewID() Role: req.Role,
_, err := database.DB.Exec(` Content: req.Content,
INSERT INTO messages (id, channel_id, role, content, model, parent_id, Model: req.Model,
participant_type, participant_id, sibling_index) ParentID: parentID,
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) SiblingIndex: siblingIdx,
`, newID, channelID, req.Role, req.Content, req.Model, parentID, ParticipantType: participantType,
participantType, participantID, siblingIdx) ParticipantID: participantID,
if err != nil { ToolCalls: models.JSONMap{},
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create message"}) Metadata: models.JSONMap{},
return
}
err = database.DB.QueryRow(`
SELECT id, channel_id, role, content, model, tokens_used, parent_id,
sibling_index, created_at
FROM messages WHERE id = ?`, newID).Scan(
&msg.ID, &msg.ChannelID, &msg.Role, &msg.Content,
&msg.Model, &msg.TokensUsed, &msg.ParentID,
&msg.SiblingIndex, &msg.CreatedAt,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to read created message"})
return
}
} else {
err := database.DB.QueryRow(`
INSERT INTO messages (channel_id, role, content, model, parent_id,
participant_type, participant_id, sibling_index)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8)
RETURNING id, channel_id, role, content, model, tokens_used, parent_id,
sibling_index, created_at
`, channelID, req.Role, req.Content, req.Model, parentID,
participantType, participantID, siblingIdx,
).Scan(
&msg.ID, &msg.ChannelID, &msg.Role, &msg.Content,
&msg.Model, &msg.TokensUsed, &msg.ParentID,
&msg.SiblingIndex, &msg.CreatedAt,
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create message"})
return
}
} }
msg.SiblingCount = getSiblingCount(channelID, msg.ParentID)
_ = updateCursor(channelID, userID, msg.ID) if err := h.stores.Messages.CreateWithCursor(c.Request.Context(), msg, userID); err != nil {
_, _ = database.DB.Exec(database.Q(`UPDATE channels SET updated_at = NOW() WHERE id = $1`), channelID) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create message"})
return
}
c.JSON(http.StatusCreated, msg) resp := messageResponse{
ID: msg.ID,
ChannelID: msg.ChannelID,
Role: msg.Role,
Content: msg.Content,
SiblingIndex: msg.SiblingIndex,
CreatedAt: msg.CreatedAt.Format("2006-01-02T15:04:05Z"),
}
if msg.Model != "" {
resp.Model = &msg.Model
}
resp.ParentID = msg.ParentID
resp.SiblingCount = getSiblingCount(channelID, msg.ParentID)
c.JSON(http.StatusCreated, resp)
} }
// ── Edit Message (create sibling) ─────────── // ── Edit Message (create sibling) ───────────
// POST /channels/:id/messages/:msgId/edit // POST /channels/:id/messages/:msgId/edit
//
// Creates a new user message as a sibling of the target (same parent_id).
// Does NOT auto-trigger completion — the frontend sends a separate request.
func (h *MessageHandler) EditMessage(c *gin.Context) { func (h *MessageHandler) EditMessage(c *gin.Context) {
var req editRequest var req editRequest
@@ -300,18 +249,14 @@ func (h *MessageHandler) EditMessage(c *gin.Context) {
channelID := c.Param("id") channelID := c.Param("id")
messageID := c.Param("msgId") messageID := c.Param("msgId")
if !userOwnsChannel(c, channelID, userID) { if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
// Load target — must exist, belong to channel, be a user message ctx := c.Request.Context()
var targetParentID *string
var targetRole string
err := database.DB.QueryRow(database.Q(`
SELECT parent_id, role FROM messages
WHERE id = $1 AND channel_id = $2 AND deleted_at IS NULL
`), messageID, channelID).Scan(&targetParentID, &targetRole)
// Load target — must exist, belong to channel, be a user message
targetParentID, targetRole, err := h.stores.Messages.GetParentAndRole(ctx, messageID, channelID)
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "message not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "message not found"})
return return
@@ -328,65 +273,39 @@ func (h *MessageHandler) EditMessage(c *gin.Context) {
// Create sibling: same parent_id as the target // Create sibling: same parent_id as the target
siblingIdx := nextSiblingIndex(channelID, targetParentID) siblingIdx := nextSiblingIndex(channelID, targetParentID)
var msg messageResponse msg := &models.Message{
if database.IsSQLite() { ChannelID: channelID,
newID := store.NewID() Role: "user",
_, err = database.DB.Exec(` Content: req.Content,
INSERT INTO messages (id, channel_id, role, content, parent_id, ParentID: targetParentID,
participant_type, participant_id, sibling_index) SiblingIndex: siblingIdx,
VALUES (?, ?, 'user', ?, ?, 'user', ?, ?) ParticipantType: "user",
`, newID, channelID, req.Content, targetParentID, userID, siblingIdx) ParticipantID: userID,
if err != nil { ToolCalls: models.JSONMap{},
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create edit"}) Metadata: models.JSONMap{},
return
}
err = database.DB.QueryRow(`
SELECT id, channel_id, role, content, model, tokens_used, parent_id,
sibling_index, created_at
FROM messages WHERE id = ?`, newID).Scan(
&msg.ID, &msg.ChannelID, &msg.Role, &msg.Content,
&msg.Model, &msg.TokensUsed, &msg.ParentID,
&msg.SiblingIndex, &msg.CreatedAt,
)
} else {
err = database.DB.QueryRow(`
INSERT INTO messages (channel_id, role, content, parent_id,
participant_type, participant_id, sibling_index)
VALUES ($1, 'user', $2, $3, 'user', $4, $5)
RETURNING id, channel_id, role, content, model, tokens_used, parent_id,
sibling_index, created_at
`, channelID, req.Content, targetParentID, userID, siblingIdx,
).Scan(
&msg.ID, &msg.ChannelID, &msg.Role, &msg.Content,
&msg.Model, &msg.TokensUsed, &msg.ParentID,
&msg.SiblingIndex, &msg.CreatedAt,
)
} }
if err != nil {
if err := h.stores.Messages.CreateWithCursor(ctx, msg, userID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create edit"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create edit"})
return return
} }
msg.SiblingCount = getSiblingCount(channelID, msg.ParentID)
// Cursor now points to the new sibling (it's a leaf — no children yet) resp := messageResponse{
_ = updateCursor(channelID, userID, msg.ID) ID: msg.ID,
_, _ = database.DB.Exec(database.Q(`UPDATE channels SET updated_at = NOW() WHERE id = $1`), channelID) ChannelID: msg.ChannelID,
Role: msg.Role,
Content: msg.Content,
ParentID: msg.ParentID,
SiblingIndex: msg.SiblingIndex,
SiblingCount: getSiblingCount(channelID, msg.ParentID),
CreatedAt: msg.CreatedAt.Format("2006-01-02T15:04:05Z"),
}
c.JSON(http.StatusCreated, msg) c.JSON(http.StatusCreated, resp)
} }
// ── Regenerate / Complete ────────────────────── // ── Regenerate / Complete ──────────────────────
// POST /channels/:id/messages/:msgId/regenerate // POST /channels/:id/messages/:msgId/regenerate
//
// Two modes depending on target role:
// - assistant message → creates a NEW assistant sibling (same parent_id).
// Context is root → target's parent. ("Give me a different response.")
// - user message → creates a child assistant response.
// Context is root → target. ("Respond to this message.")
//
// The second mode is used after editing: the edit endpoint creates a
// sibling user message, then the frontend calls regenerate on it to
// get an assistant response without duplicating the user message.
func (h *MessageHandler) Regenerate(c *gin.Context) { func (h *MessageHandler) Regenerate(c *gin.Context) {
var req regenerateRequest var req regenerateRequest
@@ -396,18 +315,14 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
channelID := c.Param("id") channelID := c.Param("id")
messageID := c.Param("msgId") messageID := c.Param("msgId")
if !userOwnsChannel(c, channelID, userID) { if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
// Load target message ctx := c.Request.Context()
var targetParentID *string
var targetRole string
err := database.DB.QueryRow(database.Q(`
SELECT parent_id, role FROM messages
WHERE id = $1 AND channel_id = $2 AND deleted_at IS NULL
`), messageID, channelID).Scan(&targetParentID, &targetRole)
// Load target message
targetParentID, targetRole, err := h.stores.Messages.GetParentAndRole(ctx, messageID, channelID)
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "message not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "message not found"})
return return
@@ -421,13 +336,7 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
return return
} }
// Determine context path and new message's parent based on target role: // Determine context path and new message's parent based on target role
//
// assistant target: context = root → target's parent (exclude target)
// new parent = target's parent (sibling of target)
//
// user target: context = root → target (include target)
// new parent = target itself (child of target)
var contextPath []PathMessage var contextPath []PathMessage
var newParentID *string var newParentID *string
@@ -486,15 +395,14 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
// Fallback: channel's stored model // Fallback: channel's stored model
if model == "" { if model == "" {
var channelModel *string channelModel, _ := h.stores.Channels.GetDefaultModel(ctx, channelID)
_ = database.DB.QueryRow(database.Q(`SELECT model FROM channels WHERE id = $1`), channelID).Scan(&channelModel)
if channelModel != nil { if channelModel != nil {
model = *channelModel model = *channelModel
} }
} }
providerCfg, providerID, model, configID, providerScope, err := comp.resolveConfig(userID, channelID, completionRequest{ providerCfg, providerID, model, configID, providerScope, err := comp.resolveConfig(userID, channelID, completionRequest{
Model: model, Model: model,
ProviderConfigID: providerConfigID, ProviderConfigID: providerConfigID,
}) })
if err != nil { if err != nil {
@@ -514,8 +422,7 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
if personaSystemPrompt != "" { if personaSystemPrompt != "" {
llmMessages = append(llmMessages, providers.Message{Role: "system", Content: personaSystemPrompt}) llmMessages = append(llmMessages, providers.Message{Role: "system", Content: personaSystemPrompt})
} else { } else {
var systemPrompt *string systemPrompt, _ := h.stores.Channels.GetSystemPrompt(ctx, channelID)
_ = database.DB.QueryRow(database.Q(`SELECT system_prompt FROM channels WHERE id = $1`), channelID).Scan(&systemPrompt)
if systemPrompt != nil && *systemPrompt != "" { if systemPrompt != nil && *systemPrompt != "" {
llmMessages = append(llmMessages, providers.Message{Role: "system", Content: *systemPrompt}) llmMessages = append(llmMessages, providers.Message{Role: "system", Content: *systemPrompt})
} }
@@ -544,14 +451,10 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
} }
// Attach tool definitions (same as normal completion) // Attach tool definitions (same as normal completion)
workspaceID, _ := h.stores.Channels.ResolveWorkspaceID(c.Request.Context(), channelID) workspaceID, _ := h.stores.Channels.ResolveWorkspaceID(ctx, channelID)
// Query channel metadata for tool context and execution context // Query channel metadata for tool context
var chType string chType, chTeamID, _ := h.stores.Channels.GetTypeAndTeamID(ctx, channelID)
var chTeamID *string
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COALESCE(type, 'direct'), team_id FROM channels WHERE id = $1
`), channelID).Scan(&chType, &chTeamID)
msgTeamID := "" msgTeamID := ""
if chTeamID != nil { if chTeamID != nil {
msgTeamID = *chTeamID msgTeamID = *chTeamID
@@ -559,7 +462,6 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
hasBrowserTools := h.hub != nil && h.hub.IsConnected(userID) hasBrowserTools := h.hub != nil && h.hub.IsConnected(userID)
if caps.ToolCalling && (tools.HasTools() || hasBrowserTools) { if caps.ToolCalling && (tools.HasTools() || hasBrowserTools) {
// v0.25.0: Build ToolContext with channel metadata
tctx := tools.ToolContext{ tctx := tools.ToolContext{
ChannelType: chType, ChannelType: chType,
WorkspaceID: workspaceID, WorkspaceID: workspaceID,
@@ -567,12 +469,11 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
PersonaID: personaID, PersonaID: personaID,
IsVisitor: isSessionAuth(c), IsVisitor: isSessionAuth(c),
} }
provReq.Tools = comp.buildToolDefs(c.Request.Context(), userID, hasBrowserTools, req.DisabledTools, tctx, personaID) provReq.Tools = comp.buildToolDefs(ctx, userID, hasBrowserTools, req.DisabledTools, tctx, personaID)
} }
// ── Stream the response (shared loop handles tools, reasoning, SSE) ── // ── Stream the response ──
// Apply provider-specific request hooks (v0.22.1)
if hooks := providers.GetHooks(providerID); hooks != nil { if hooks := providers.GetHooks(providerID); hooks != nil {
hooks.PreRequest(providerCfg, &provReq) hooks.PreRequest(providerCfg, &provReq)
} }
@@ -583,47 +484,34 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
if result.Content != "" { if result.Content != "" {
siblingIdx := nextSiblingIndex(channelID, newParentID) siblingIdx := nextSiblingIndex(channelID, newParentID)
// Match persistMessage pattern: nil interface{} → SQL NULL, var toolCalls models.JSONMap
// non-nil string → valid JSONB. A nil json.RawMessage ([]byte)
// is sent by pq as empty bytes, not NULL, causing "invalid input
// syntax for type json".
var tcVal interface{}
if len(result.ToolActivity) > 0 { if len(result.ToolActivity) > 0 {
b, _ := json.Marshal(result.ToolActivity) b, _ := json.Marshal(result.ToolActivity)
tcVal = string(b) _ = json.Unmarshal(b, &toolCalls)
}
if toolCalls == nil {
toolCalls = models.JSONMap{}
} }
var newID string msg := &models.Message{
if database.IsSQLite() { ChannelID: channelID,
newID = store.NewID() Role: "assistant",
_, err := database.DB.Exec(` Content: result.Content,
INSERT INTO messages (id, channel_id, role, content, model, tool_calls, Model: model,
parent_id, participant_type, participant_id, sibling_index) ToolCalls: toolCalls,
VALUES (?, ?, 'assistant', ?, ?, ?, ?, 'model', ?, ?) Metadata: models.JSONMap{},
`, newID, channelID, result.Content, model, tcVal, newParentID, model, siblingIdx) ParentID: newParentID,
if err != nil { SiblingIndex: siblingIdx,
log.Printf("Failed to persist regenerated message: %v", err) ParticipantType: "model",
newID = "" ParticipantID: model,
} }
if err := h.stores.Messages.CreateWithCursor(ctx, msg, userID); err != nil {
log.Printf("Failed to persist regenerated message: %v", err)
} else { } else {
err := database.DB.QueryRow(`
INSERT INTO messages (channel_id, role, content, model, tool_calls,
parent_id, participant_type, participant_id, sibling_index)
VALUES ($1, 'assistant', $2, $3, $4, $5, 'model', $6, $7)
RETURNING id
`, channelID, result.Content, model, tcVal, newParentID, model, siblingIdx).Scan(&newID)
if err != nil {
log.Printf("Failed to persist regenerated message: %v", err)
newID = ""
}
}
if newID != "" {
_ = updateCursor(channelID, userID, newID)
flusher, _ := c.Writer.(http.Flusher) flusher, _ := c.Writer.(http.Flusher)
msgJSON, _ := json.Marshal(gin.H{ msgJSON, _ := json.Marshal(gin.H{
"id": newID, "id": msg.ID,
"parent_id": newParentID, "parent_id": newParentID,
"sibling_index": siblingIdx, "sibling_index": siblingIdx,
"sibling_count": getSiblingCount(channelID, newParentID), "sibling_count": getSiblingCount(channelID, newParentID),
@@ -633,8 +521,6 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
flusher.Flush() flusher.Flush()
} }
} }
_, _ = database.DB.Exec(database.Q(`UPDATE channels SET updated_at = NOW() WHERE id = $1`), channelID)
} }
// Log usage for regeneration // Log usage for regeneration
@@ -645,9 +531,6 @@ func (h *MessageHandler) Regenerate(c *gin.Context) {
// ── Switch Branch (update cursor) ─────────── // ── Switch Branch (update cursor) ───────────
// PUT /channels/:id/cursor // PUT /channels/:id/cursor
//
// Navigates to a different branch. If the target message has children,
// follows the first-child chain to find the leaf.
func (h *MessageHandler) UpdateCursor(c *gin.Context) { func (h *MessageHandler) UpdateCursor(c *gin.Context) {
var req cursorRequest var req cursorRequest
@@ -659,17 +542,15 @@ func (h *MessageHandler) UpdateCursor(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
if !userOwnsChannel(c, channelID, userID) { if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
// Verify the target message belongs to this channel ctx := c.Request.Context()
var msgChannelID string
err := database.DB.QueryRow(database.Q(`
SELECT channel_id FROM messages
WHERE id = $1 AND deleted_at IS NULL
`), req.ActiveLeafID).Scan(&msgChannelID)
// Verify the target message belongs to this channel and is not deleted.
// GetParentAndRole does: WHERE id=$1 AND channel_id=$2 AND deleted_at IS NULL
_, _, err := h.stores.Messages.GetParentAndRole(ctx, req.ActiveLeafID, channelID)
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "message not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "message not found"})
return return
@@ -678,10 +559,6 @@ func (h *MessageHandler) UpdateCursor(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify message"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify message"})
return return
} }
if msgChannelID != channelID {
c.JSON(http.StatusForbidden, gin.H{"error": "message does not belong to this channel"})
return
}
// Walk down to the leaf from the target // Walk down to the leaf from the target
leafID, err := findLeafFromMessage(req.ActiveLeafID) leafID, err := findLeafFromMessage(req.ActiveLeafID)
@@ -707,16 +584,13 @@ func (h *MessageHandler) UpdateCursor(c *gin.Context) {
// ── List Siblings ─────────────────────────── // ── List Siblings ───────────────────────────
// GET /channels/:id/messages/:msgId/siblings // GET /channels/:id/messages/:msgId/siblings
//
// Returns all siblings of a message (messages sharing the same parent_id).
// Used by the branch navigator to populate the arrow navigation.
func (h *MessageHandler) ListSiblings(c *gin.Context) { func (h *MessageHandler) ListSiblings(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
channelID := c.Param("id") channelID := c.Param("id")
messageID := c.Param("msgId") messageID := c.Param("msgId")
if !userOwnsChannel(c, channelID, userID) { if !userCanAccessChannel(c, h.stores, channelID, userID) {
return return
} }
@@ -733,36 +607,34 @@ func (h *MessageHandler) ListSiblings(c *gin.Context) {
}) })
} }
// ── Ownership Check ───────────────────────── // ── Ownership / Access Check ─────────────────
func userOwnsChannel(c *gin.Context, channelID, userID string) bool { // userCanAccessChannel verifies channel ownership or participation,
var ownerID string // writing an error response if denied. Returns true if access is allowed.
err := database.DB.QueryRow( func userCanAccessChannel(c *gin.Context, stores store.Stores, channelID, userID string) bool {
database.Q(`SELECT user_id FROM channels WHERE id = $1`), channelID, ok, err := stores.Channels.UserCanAccess(c.Request.Context(), channelID, userID)
).Scan(&ownerID)
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return false
}
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify channel ownership"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify channel access"})
return false return false
} }
if ownerID == userID { if !ok {
return true c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"})
return false
} }
return true
// v0.23.2: Also allow if user is a participant (multi-user DMs, channels) }
var exists bool
_ = database.DB.QueryRow( // userOwnsChannel is the legacy wrapper used by other handlers (completion.go).
database.Q(`SELECT EXISTS(SELECT 1 FROM channel_participants WHERE channel_id = $1 AND participant_type = 'user' AND participant_id = $2)`), // Delegates to userCanAccessChannel using the treepath global stores.
channelID, userID, func userOwnsChannel(c *gin.Context, channelID, userID string) bool {
).Scan(&exists) ok, err := treepath.Stores.Channels.UserCanAccess(c.Request.Context(), channelID, userID)
if exists { if err != nil {
return true c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify channel access"})
} return false
}
c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"}) if !ok {
return false c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"})
return false
}
return true
} }

View File

@@ -7,7 +7,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/notelinks" "git.gobha.me/xcaliber/chat-switchboard/notelinks"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
@@ -405,69 +404,30 @@ func (h *NoteHandler) Search(c *gin.Context) {
limit = 20 limit = 20
} }
// Use Postgres full-text search when available, LIKE fallback on SQLite // SearchKeyword handles dialect internally: PG uses ts_rank/ts_headline, SQLite uses LIKE.
if database.IsPostgres() { storeResults, err := h.stores.Notes.SearchKeyword(c.Request.Context(), userID, q, limit)
h.searchPostgres(c, userID, q, limit)
return
}
// SQLite: use store's LIKE-based search
notes, _, err := h.stores.Notes.Search(c.Request.Context(), userID, q, store.ListOptions{Limit: limit})
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "search failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "search failed"})
return return
} }
results := make([]searchResult, 0, len(notes)) results := make([]searchResult, 0, len(storeResults))
for _, n := range notes { for _, sr := range storeResults {
results = append(results, searchResult{ tags := sr.Tags
noteListItem: toNoteListItem(n),
Rank: 1.0,
Headline: "",
})
}
c.JSON(http.StatusOK, gin.H{
"data": results,
"query": q,
"total": len(results),
})
}
// searchPostgres uses Postgres full-text search with ts_rank and ts_headline.
func (h *NoteHandler) searchPostgres(c *gin.Context, userID, q string, limit int) {
rows, err := database.DB.Query(`
SELECT id, title, folder_path, tags, LEFT(content, 200),
created_at::text, updated_at::text,
ts_rank(search_vector, plainto_tsquery('english', $2)) AS rank,
ts_headline('english', content, plainto_tsquery('english', $2),
'MaxWords=40, MinWords=20, StartSel=**, StopSel=**') AS headline
FROM notes
WHERE user_id = $1
AND search_vector @@ plainto_tsquery('english', $2)
ORDER BY rank DESC
LIMIT $3
`, userID, q, limit)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "search failed"})
return
}
defer rows.Close()
// Import pq at call site to avoid pulling it in for SQLite builds
results := make([]searchResult, 0)
for rows.Next() {
var r searchResult
var tags []string
if err := rows.Scan(&r.ID, &r.Title, &r.FolderPath, pgScanStringArray(&tags), &r.Preview,
&r.CreatedAt, &r.UpdatedAt, &r.Rank, &r.Headline); err != nil {
continue
}
if tags == nil { if tags == nil {
tags = []string{} tags = []string{}
} }
r.Tags = tags results = append(results, searchResult{
results = append(results, r) noteListItem: noteListItem{
ID: sr.ID,
Title: sr.Title,
FolderPath: sr.FolderPath,
Tags: tags,
Preview: sr.Excerpt,
},
Rank: sr.Rank,
Headline: sr.Headline,
})
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
@@ -483,29 +443,19 @@ func (h *NoteHandler) searchPostgres(c *gin.Context, userID, q string, limit int
func (h *NoteHandler) ListFolders(c *gin.Context) { func (h *NoteHandler) ListFolders(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
rows, err := database.DB.Query(database.Q(` storeFolders, err := h.stores.Notes.ListFolders(c.Request.Context(), userID)
SELECT DISTINCT folder_path, COUNT(*) AS count
FROM notes WHERE user_id = $1
GROUP BY folder_path
ORDER BY folder_path
`), userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list folders"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list folders"})
return return
} }
defer rows.Close()
type folderInfo struct { type folderInfo struct {
Path string `json:"path"` Path string `json:"path"`
Count int `json:"count"` Count int `json:"count"`
} }
folders := make([]folderInfo, 0) folders := make([]folderInfo, 0, len(storeFolders))
for rows.Next() { for _, f := range storeFolders {
var f folderInfo folders = append(folders, folderInfo{Path: f.Path, Count: f.Count})
if err := rows.Scan(&f.Path, &f.Count); err != nil {
continue
}
folders = append(folders, f)
} }
c.JSON(http.StatusOK, gin.H{"folders": folders}) c.JSON(http.StatusOK, gin.H{"folders": folders})

View File

@@ -6,7 +6,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
@@ -260,17 +259,10 @@ func (h *ParticipantHandler) Remove(c *gin.Context) {
} }
// v0.23.2: DM guard — cannot drop below 2 human participants // v0.23.2: DM guard — cannot drop below 2 human participants
var channelType string channelType, _, _ := h.stores.Channels.GetTypeAndTeamID(c.Request.Context(), channelID)
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COALESCE(type, 'direct') FROM channels WHERE id = $1
`), channelID).Scan(&channelType)
if channelType == "dm" && p.ParticipantType == "user" { if channelType == "dm" && p.ParticipantType == "user" {
var userCount int userCount, _ := h.stores.Channels.CountParticipantsByType(c.Request.Context(), channelID, "user")
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COUNT(*) FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'user'
`), channelID).Scan(&userCount)
if userCount <= 2 { if userCount <= 2 {
c.JSON(http.StatusConflict, gin.H{"error": "DMs require at least 2 participants"}) c.JSON(http.StatusConflict, gin.H{"error": "DMs require at least 2 participants"})
return return
@@ -279,11 +271,7 @@ func (h *ParticipantHandler) Remove(c *gin.Context) {
// v0.23.2: Group guard — cannot remove the last persona // v0.23.2: Group guard — cannot remove the last persona
if channelType == "group" && p.ParticipantType == "persona" { if channelType == "group" && p.ParticipantType == "persona" {
var personaCount int personaCount, _ := h.stores.Channels.CountParticipantsByType(c.Request.Context(), channelID, "persona")
_ = database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT COUNT(*) FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'persona'
`), channelID).Scan(&personaCount)
if personaCount <= 1 { if personaCount <= 1 {
c.JSON(http.StatusConflict, gin.H{"error": "group chats require at least 1 persona"}) c.JSON(http.StatusConflict, gin.H{"error": "group chats require at least 1 persona"})
return return

View File

@@ -2,64 +2,40 @@ package handlers
// persona_groups.go — Persona group (roster template) CRUD (v0.23.2) // persona_groups.go — Persona group (roster template) CRUD (v0.23.2)
// //
// Persona groups are saved collections of personas used as templates // v0.29.0: Raw SQL replaced with PersonaGroupStore methods.
// for creating group chats. Each member has an is_leader flag.
//
// Routes:
// GET /api/v1/persona-groups
// POST /api/v1/persona-groups
// GET /api/v1/persona-groups/:id
// PUT /api/v1/persona-groups/:id
// DELETE /api/v1/persona-groups/:id
// POST /api/v1/persona-groups/:id/members
// DELETE /api/v1/persona-groups/:id/members/:memberId
import ( import (
"database/sql"
"net/http" "net/http"
"strings" "strings"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
type PersonaGroupHandler struct{} type PersonaGroupHandler struct {
stores store.Stores
}
func NewPersonaGroupHandler() *PersonaGroupHandler { return &PersonaGroupHandler{} } func NewPersonaGroupHandler(s store.Stores) *PersonaGroupHandler {
return &PersonaGroupHandler{stores: s}
}
// ── List ──────────────────────────────────────── // ── List ────────────────────────────────────────
func (h *PersonaGroupHandler) List(c *gin.Context) { func (h *PersonaGroupHandler) List(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(` groups, err := h.stores.PersonaGroups.List(c.Request.Context(), userID)
SELECT id, name, description, owner_id, scope, team_id, created_at, updated_at
FROM persona_groups
WHERE owner_id = $1
ORDER BY name
`), userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list groups"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list groups"})
return return
} }
defer rows.Close()
groups := []models.PersonaGroup{}
for rows.Next() {
var g models.PersonaGroup
if err := rows.Scan(&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, database.ST(&g.CreatedAt), database.ST(&g.UpdatedAt)); err != nil {
continue
}
groups = append(groups, g)
}
// Load members for each group
for i := range groups { for i := range groups {
groups[i].Members = h.loadMembers(c, groups[i].ID) members, _ := h.stores.PersonaGroups.ListMembers(c.Request.Context(), groups[i].ID)
groups[i].Members = members
} }
c.JSON(http.StatusOK, gin.H{"data": groups}) c.JSON(http.StatusOK, gin.H{"data": groups})
@@ -71,25 +47,22 @@ func (h *PersonaGroupHandler) Get(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
id := c.Param("id") id := c.Param("id")
var g models.PersonaGroup g, err := h.stores.PersonaGroups.Get(c.Request.Context(), id, userID)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT id, name, description, owner_id, scope, team_id, created_at, updated_at
FROM persona_groups WHERE id = $1 AND owner_id = $2
`), id, userID).Scan(&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, database.ST(&g.CreatedAt), database.ST(&g.UpdatedAt))
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "group not found"})
return
}
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to get group"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to get group"})
return return
} }
if g == nil {
g.Members = h.loadMembers(c, g.ID) c.JSON(http.StatusNotFound, gin.H{"error": "group not found"})
if g.Members == nil { return
g.Members = []models.PersonaGroupMember{}
} }
members, _ := h.stores.PersonaGroups.ListMembers(c.Request.Context(), g.ID)
if members == nil {
members = []models.PersonaGroupMember{}
}
g.Members = members
c.JSON(http.StatusOK, g) c.JSON(http.StatusOK, g)
} }
@@ -107,34 +80,16 @@ func (h *PersonaGroupHandler) Create(c *gin.Context) {
return return
} }
var g models.PersonaGroup g := &models.PersonaGroup{
if database.IsSQLite() { Name: strings.TrimSpace(req.Name),
id := store.NewID() Description: req.Description,
_, err := database.DB.ExecContext(c.Request.Context(), ` OwnerID: userID,
INSERT INTO persona_groups (id, name, description, owner_id, scope) Scope: "personal",
VALUES (?, ?, ?, ?, 'personal') }
`, id, strings.TrimSpace(req.Name), req.Description, userID)
if err != nil { if err := h.stores.PersonaGroups.Create(c.Request.Context(), g); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create group"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create group"})
return return
}
database.DB.QueryRowContext(c.Request.Context(), `
SELECT id, name, description, owner_id, scope, team_id, created_at, updated_at
FROM persona_groups WHERE id = ?
`, id).Scan(&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, database.ST(&g.CreatedAt), database.ST(&g.UpdatedAt))
} else {
err := database.DB.QueryRowContext(c.Request.Context(), `
INSERT INTO persona_groups (name, description, owner_id, scope)
VALUES ($1, $2, $3, 'personal')
RETURNING id, name, description, owner_id, scope, team_id, created_at, updated_at
`, strings.TrimSpace(req.Name), req.Description, userID).Scan(
&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, database.ST(&g.CreatedAt), database.ST(&g.UpdatedAt))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create group"})
return
}
} }
g.Members = []models.PersonaGroupMember{} g.Members = []models.PersonaGroupMember{}
@@ -157,11 +112,8 @@ func (h *PersonaGroupHandler) Update(c *gin.Context) {
} }
// Verify ownership // Verify ownership
var ownerID string ownerID, err := h.stores.PersonaGroups.GetOwnerID(c.Request.Context(), id)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(` if err != nil || ownerID == "" {
SELECT owner_id FROM persona_groups WHERE id = $1
`), id).Scan(&ownerID)
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "group not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "group not found"})
return return
} }
@@ -170,15 +122,15 @@ func (h *PersonaGroupHandler) Update(c *gin.Context) {
return return
} }
fields := map[string]interface{}{}
if req.Name != nil { if req.Name != nil {
database.DB.ExecContext(c.Request.Context(), database.Q(` fields["name"] = strings.TrimSpace(*req.Name)
UPDATE persona_groups SET name = $1, updated_at = NOW() WHERE id = $2
`), strings.TrimSpace(*req.Name), id)
} }
if req.Description != nil { if req.Description != nil {
database.DB.ExecContext(c.Request.Context(), database.Q(` fields["description"] = *req.Description
UPDATE persona_groups SET description = $1, updated_at = NOW() WHERE id = $2 }
`), *req.Description, id) if len(fields) > 0 {
_ = h.stores.PersonaGroups.Update(c.Request.Context(), id, fields)
} }
c.JSON(http.StatusOK, gin.H{"ok": true}) c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -190,14 +142,11 @@ func (h *PersonaGroupHandler) Delete(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
id := c.Param("id") id := c.Param("id")
result, err := database.DB.ExecContext(c.Request.Context(), database.Q(` n, err := h.stores.PersonaGroups.Delete(c.Request.Context(), id, userID)
DELETE FROM persona_groups WHERE id = $1 AND owner_id = $2
`), id, userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "delete failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "delete failed"})
return return
} }
n, _ := result.RowsAffected()
if n == 0 { if n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "group not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "group not found"})
return return
@@ -221,37 +170,13 @@ func (h *PersonaGroupHandler) AddMember(c *gin.Context) {
} }
// Verify ownership // Verify ownership
var ownerID string ownerID, err := h.stores.PersonaGroups.GetOwnerID(c.Request.Context(), groupID)
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT owner_id FROM persona_groups WHERE id = $1
`), groupID).Scan(&ownerID)
if err != nil || ownerID != userID { if err != nil || ownerID != userID {
c.JSON(http.StatusForbidden, gin.H{"error": "not your group"}) c.JSON(http.StatusForbidden, gin.H{"error": "not your group"})
return return
} }
// If setting as leader, clear existing leader if err := h.stores.PersonaGroups.AddMember(c.Request.Context(), groupID, req.PersonaID, req.IsLeader); err != nil {
if req.IsLeader {
database.DB.ExecContext(c.Request.Context(), database.Q(`
UPDATE persona_group_members SET is_leader = false WHERE group_id = $1
`), groupID)
}
if database.IsSQLite() {
id := store.NewID()
_, err = database.DB.ExecContext(c.Request.Context(), `
INSERT INTO persona_group_members (id, group_id, persona_id, is_leader)
VALUES (?, ?, ?, ?)
ON CONFLICT (group_id, persona_id) DO UPDATE SET is_leader = excluded.is_leader
`, id, groupID, req.PersonaID, req.IsLeader)
} else {
_, err = database.DB.ExecContext(c.Request.Context(), `
INSERT INTO persona_group_members (group_id, persona_id, is_leader)
VALUES ($1, $2, $3)
ON CONFLICT (group_id, persona_id) DO UPDATE SET is_leader = EXCLUDED.is_leader
`, groupID, req.PersonaID, req.IsLeader)
}
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to add member"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to add member"})
return return
} }
@@ -266,49 +191,12 @@ func (h *PersonaGroupHandler) RemoveMember(c *gin.Context) {
groupID := c.Param("id") groupID := c.Param("id")
memberID := c.Param("memberId") memberID := c.Param("memberId")
// Verify ownership ownerID, err := h.stores.PersonaGroups.GetOwnerID(c.Request.Context(), groupID)
var ownerID string
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT owner_id FROM persona_groups WHERE id = $1
`), groupID).Scan(&ownerID)
if err != nil || ownerID != userID { if err != nil || ownerID != userID {
c.JSON(http.StatusForbidden, gin.H{"error": "not your group"}) c.JSON(http.StatusForbidden, gin.H{"error": "not your group"})
return return
} }
database.DB.ExecContext(c.Request.Context(), database.Q(` _ = h.stores.PersonaGroups.RemoveMember(c.Request.Context(), memberID, groupID)
DELETE FROM persona_group_members WHERE id = $1 AND group_id = $2
`), memberID, groupID)
c.JSON(http.StatusOK, gin.H{"ok": true}) c.JSON(http.StatusOK, gin.H{"ok": true})
} }
// ── Helpers ─────────────────────────────────────
func (h *PersonaGroupHandler) loadMembers(c *gin.Context, groupID string) []models.PersonaGroupMember {
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(`
SELECT pgm.id, pgm.group_id, pgm.persona_id, pgm.is_leader, pgm.sort_order,
COALESCE(p.name, '') AS persona_name,
COALESCE(p.handle, '') AS persona_handle,
COALESCE(p.avatar, '') AS persona_avatar
FROM persona_group_members pgm
LEFT JOIN personas p ON p.id = pgm.persona_id
WHERE pgm.group_id = $1
ORDER BY pgm.is_leader DESC, pgm.sort_order, pgm.id
`), groupID)
if err != nil {
return []models.PersonaGroupMember{}
}
defer rows.Close()
members := []models.PersonaGroupMember{}
for rows.Next() {
var m models.PersonaGroupMember
if err := rows.Scan(&m.ID, &m.GroupID, &m.PersonaID, &m.IsLeader, &m.SortOrder,
&m.PersonaName, &m.PersonaHandle, &m.PersonaAvatar); err != nil {
continue
}
members = append(members, m)
}
return members
}

View File

@@ -450,7 +450,7 @@ func (h *PersonaHandler) UploadUserPersonaAvatar(c *gin.Context) {
} }
// Delegate to the shared avatar handler // Delegate to the shared avatar handler
UploadPersonaAvatar(c) UploadPersonaAvatar(h.stores.Personas, c)
} }
// DeleteUserPersonaAvatar deletes an avatar for a personal persona, // DeleteUserPersonaAvatar deletes an avatar for a personal persona,
@@ -471,7 +471,7 @@ func (h *PersonaHandler) DeleteUserPersonaAvatar(c *gin.Context) {
} }
// Delegate to the shared avatar handler // Delegate to the shared avatar handler
DeletePersonaAvatar(c) DeletePersonaAvatar(h.stores.Personas, c)
} }
// ── Team-Scoped Helpers ───────────────────── // ── Team-Scoped Helpers ─────────────────────
@@ -556,7 +556,7 @@ func (h *PersonaHandler) UploadTeamPersonaAvatar(c *gin.Context) {
if p := h.requireTeamPersona(c); p == nil { if p := h.requireTeamPersona(c); p == nil {
return return
} }
UploadPersonaAvatar(c) UploadPersonaAvatar(h.stores.Personas, c)
} }
// DeleteTeamPersonaAvatar deletes an avatar for a team persona, // DeleteTeamPersonaAvatar deletes an avatar for a team persona,
@@ -565,5 +565,5 @@ func (h *PersonaHandler) DeleteTeamPersonaAvatar(c *gin.Context) {
if p := h.requireTeamPersona(c); p == nil { if p := h.requireTeamPersona(c); p == nil {
return return
} }
DeletePersonaAvatar(c) DeletePersonaAvatar(h.stores.Personas, c)
} }

View File

@@ -5,6 +5,8 @@ package handlers
// Clients POST /api/v1/presence/heartbeat every 30s while active. // Clients POST /api/v1/presence/heartbeat every 30s while active.
// GET /api/v1/presence?users=id1,id2 returns current status. // GET /api/v1/presence?users=id1,id2 returns current status.
// Online = last_seen within 90s. // Online = last_seen within 90s.
//
// v0.29.0: Raw SQL replaced with PresenceStore + UserStore methods.
import ( import (
"net/http" "net/http"
@@ -13,21 +15,23 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
const presenceOnlineThreshold = 90 * time.Second const presenceOnlineThreshold = 90 * time.Second
type PresenceHandler struct {
stores store.Stores
}
func NewPresenceHandler(s store.Stores) *PresenceHandler {
return &PresenceHandler{stores: s}
}
// PresenceHeartbeat upserts the calling user's last_seen timestamp. // PresenceHeartbeat upserts the calling user's last_seen timestamp.
func PresenceHeartbeat(c *gin.Context) { func (h *PresenceHandler) Heartbeat(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
_, err := database.DB.ExecContext(c.Request.Context(), database.Q(` if err := h.stores.Presence.Heartbeat(c.Request.Context(), userID); err != nil {
INSERT INTO user_presence (user_id, last_seen, status)
VALUES ($1, NOW(), 'online')
ON CONFLICT (user_id) DO UPDATE
SET last_seen = NOW(), status = 'online'
`), userID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "presence update failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "presence update failed"})
return return
} }
@@ -36,7 +40,7 @@ func PresenceHeartbeat(c *gin.Context) {
// PresenceQuery returns online/offline status for a list of user IDs. // PresenceQuery returns online/offline status for a list of user IDs.
// Query param: ?users=uuid1,uuid2,... // Query param: ?users=uuid1,uuid2,...
func PresenceQuery(c *gin.Context) { func (h *PresenceHandler) Query(c *gin.Context) {
raw := c.Query("users") raw := c.Query("users")
if raw == "" { if raw == "" {
c.JSON(http.StatusOK, gin.H{"presence": map[string]string{}}) c.JSON(http.StatusOK, gin.H{"presence": map[string]string{}})
@@ -47,71 +51,34 @@ func PresenceQuery(c *gin.Context) {
ids = ids[:100] ids = ids[:100]
} }
threshold := time.Now().Add(-presenceOnlineThreshold) // Trim whitespace
cleaned := make([]string, 0, len(ids))
result := make(map[string]string, len(ids))
for _, id := range ids { for _, id := range ids {
id = strings.TrimSpace(id) id = strings.TrimSpace(id)
if id == "" { if id != "" {
continue cleaned = append(cleaned, id)
}
var lastSeen time.Time
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT last_seen FROM user_presence WHERE user_id = $1
`), id).Scan(&lastSeen)
if err != nil || lastSeen.Before(threshold) {
result[id] = "offline"
} else {
result[id] = "online"
} }
} }
threshold := time.Now().Add(-presenceOnlineThreshold)
result, err := h.stores.Presence.GetStatuses(c.Request.Context(), cleaned, threshold)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "presence query failed"})
return
}
c.JSON(http.StatusOK, gin.H{"presence": result}) c.JSON(http.StatusOK, gin.H{"presence": result})
} }
// SearchUsers returns a lightweight list of approved users matching a query. // SearchUsers returns a lightweight list of approved users matching a query.
// GET /api/v1/users/search?q=alice — matches against username and display_name. // GET /api/v1/users/search?q=alice
// Returns at most 20 results. Excludes the calling user. func (h *PresenceHandler) SearchUsers(c *gin.Context) {
func SearchUsers(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
q := strings.TrimSpace(c.Query("q")) q := strings.TrimSpace(c.Query("q"))
query := database.Q(` results, err := h.stores.Users.SearchActive(c.Request.Context(), userID, q)
SELECT id, username, COALESCE(display_name, '') AS display_name, COALESCE(handle, '') AS handle
FROM users
WHERE is_active = true AND id != $1
`)
args := []interface{}{userID}
if q != "" {
query += database.Q(` AND (LOWER(username) LIKE $2 OR LOWER(display_name) LIKE $3 OR LOWER(handle) LIKE $4)`)
pattern := "%" + strings.ToLower(q) + "%"
args = append(args, pattern, pattern, pattern)
}
query += ` ORDER BY username LIMIT 20`
rows, err := database.DB.QueryContext(c.Request.Context(), query, args...)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "search failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "search failed"})
return return
} }
defer rows.Close()
type userResult struct {
ID string `json:"id"`
Username string `json:"username"`
DisplayName string `json:"display_name"`
Handle string `json:"handle"`
}
results := []userResult{}
for rows.Next() {
var u userResult
if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Handle); err != nil {
continue
}
results = append(results, u)
}
c.JSON(http.StatusOK, gin.H{"users": results}) c.JSON(http.StatusOK, gin.H{"users": results})
} }

View File

@@ -47,7 +47,7 @@ func setupProfileHarness(t *testing.T) *profileHarness {
protected := api.Group("") protected := api.Group("")
protected.Use(middleware.Auth(cfg, stores.Users, userCache)) protected.Use(middleware.Auth(cfg, stores.Users, userCache))
settings := NewSettingsHandler(nil) settings := NewSettingsHandler(stores, nil)
protected.GET("/profile", settings.GetProfile) protected.GET("/profile", settings.GetProfile)
protected.PUT("/profile", settings.UpdateProfile) protected.PUT("/profile", settings.UpdateProfile)
protected.POST("/profile/password", settings.ChangePassword) protected.POST("/profile/password", settings.ChangePassword)

View File

@@ -6,7 +6,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
@@ -381,24 +380,11 @@ func (h *ProjectHandler) AdminList(c *gin.Context) {
ctx := c.Request.Context() ctx := c.Request.Context()
includeArchived := c.Query("include_archived") == "true" includeArchived := c.Query("include_archived") == "true"
// Admin sees everything — query with no user scope projects, err := h.stores.Projects.AdminList(ctx, includeArchived)
rows, err := database.DB.QueryContext(ctx, database.Q(`
SELECT p.id, p.name, p.description, p.scope,
p.owner_id, p.team_id, p.is_archived,
p.created_at, p.updated_at,
(SELECT COUNT(*) FROM project_channels WHERE project_id = p.id),
(SELECT COUNT(*) FROM project_knowledge_bases WHERE project_id = p.id),
(SELECT COUNT(*) FROM project_notes WHERE project_id = p.id),
u.username
FROM projects p
LEFT JOIN users u ON u.id = p.owner_id
WHERE ($1 OR p.is_archived = false)
ORDER BY p.updated_at DESC`), includeArchived)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list projects"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list projects"})
return return
} }
defer rows.Close()
type adminProject struct { type adminProject struct {
ID string `json:"id"` ID string `json:"id"`
@@ -416,25 +402,25 @@ func (h *ProjectHandler) AdminList(c *gin.Context) {
OwnerName string `json:"owner_name"` OwnerName string `json:"owner_name"`
} }
var projects []adminProject result := make([]adminProject, 0, len(projects))
for rows.Next() { for _, p := range projects {
var p adminProject result = append(result, adminProject{
if err := rows.Scan( ID: p.ID,
&p.ID, &p.Name, &p.Description, &p.Scope, Name: p.Name,
&p.OwnerID, &p.TeamID, &p.IsArchived, Description: p.Description,
&p.CreatedAt, &p.UpdatedAt, Scope: p.Scope,
&p.ChannelCount, &p.KBCount, &p.NoteCount, OwnerID: p.OwnerID,
&p.OwnerName, TeamID: p.TeamID,
); err != nil { IsArchived: p.IsArchived,
c.JSON(http.StatusInternalServerError, gin.H{"error": "scan error"}) CreatedAt: p.CreatedAt.Format("2006-01-02T15:04:05Z"),
return UpdatedAt: p.UpdatedAt.Format("2006-01-02T15:04:05Z"),
} ChannelCount: p.ChannelCount,
projects = append(projects, p) KBCount: p.KBCount,
NoteCount: p.NoteCount,
OwnerName: p.OwnerName,
})
} }
if projects == nil { c.JSON(http.StatusOK, gin.H{"data": result})
projects = []adminProject{}
}
c.JSON(http.StatusOK, gin.H{"data": projects})
} }
// ── Auth Helper ───────────────────────────── // ── Auth Helper ─────────────────────────────

View File

@@ -6,6 +6,8 @@
// //
// The CompletionHandler methods (resolveConfig, buildToolDefs) remain as // The CompletionHandler methods (resolveConfig, buildToolDefs) remain as
// thin wrappers for backward compat — existing callers are unchanged. // thin wrappers for backward compat — existing callers are unchanged.
//
// v0.29.0-cs7a: Replaced raw SQL with store methods.
package handlers package handlers
import ( import (
@@ -16,7 +18,6 @@ import (
"log" "log"
"git.gobha.me/xcaliber/chat-switchboard/crypto" "git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/providers" "git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
"git.gobha.me/xcaliber/chat-switchboard/tools" "git.gobha.me/xcaliber/chat-switchboard/tools"
@@ -41,17 +42,16 @@ type ProviderResolution struct {
// //
// The vault is used to decrypt API keys. Pass nil for unencrypted fallback. // The vault is used to decrypt API keys. Pass nil for unencrypted fallback.
func ResolveProviderConfig( func ResolveProviderConfig(
stores store.Stores,
vault *crypto.KeyResolver, vault *crypto.KeyResolver,
userID, channelID, providerConfigID, modelID string, userID, channelID, providerConfigID, modelID string,
) (ProviderResolution, error) { ) (ProviderResolution, error) {
ctx := context.Background()
configID := providerConfigID configID := providerConfigID
// 2. Config from channel // 2. Config from channel
if configID == "" && channelID != "" { if configID == "" && channelID != "" {
var channelConfigID *string channelConfigID, err := stores.Channels.GetProviderConfigID(ctx, channelID)
err := database.DB.QueryRow(
database.Q(`SELECT provider_config_id FROM channels WHERE id = $1`), channelID,
).Scan(&channelConfigID)
if err == nil && channelConfigID != nil { if err == nil && channelConfigID != nil {
configID = *channelConfigID configID = *channelConfigID
} }
@@ -59,47 +59,15 @@ func ResolveProviderConfig(
// 3. User's first active config (personal first, then global — excludes team providers) // 3. User's first active config (personal first, then global — excludes team providers)
if configID == "" { if configID == "" {
err := database.DB.QueryRow(database.Q(` var err error
SELECT id FROM provider_configs configID, err = stores.Providers.FindFirstForUser(ctx, userID)
WHERE is_active = true AND (
(scope = 'personal' AND owner_id = $1)
OR scope = 'global'
)
ORDER BY scope ASC, created_at ASC
LIMIT 1
`), userID).Scan(&configID)
if err != nil { if err != nil {
return ProviderResolution{}, fmt.Errorf("no API config found — add one at /api-configs") return ProviderResolution{}, fmt.Errorf("no API config found — add one at /api-configs")
} }
} }
// Load the config — allow personal, global, OR team configs the user belongs to // Load the config — allow personal, global, OR team configs the user belongs to
var providerID, endpoint string cfg, err := stores.Providers.LoadAccessible(ctx, configID, userID)
var providerScope string
var modelDefault *string
var apiKeyEnc, keyNonce []byte
var keyScope string
var customHeadersJSON, providerSettingsJSON []byte
var proxyMode string
var proxyURL *string
// $2/userID appears twice in the query; Postgres reuses positional params,
// SQLite needs each ? bound separately.
configArgs := []interface{}{configID, userID}
if database.IsSQLite() {
configArgs = append(configArgs, userID)
}
err := database.DB.QueryRow(database.Q(`
SELECT provider, endpoint, scope, api_key_enc, key_nonce, key_scope,
model_default, headers, settings, COALESCE(proxy_mode, 'system'), proxy_url
FROM provider_configs
WHERE id = $1 AND is_active = true
AND (scope = 'global'
OR (scope = 'personal' AND owner_id = $2)
OR (scope = 'team' AND owner_id IN (SELECT team_id FROM team_members WHERE user_id = $2)))
`), configArgs...).Scan(&providerID, &endpoint, &providerScope, &apiKeyEnc, &keyNonce, &keyScope,
&modelDefault, &customHeadersJSON, &providerSettingsJSON, &proxyMode, &proxyURL)
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return ProviderResolution{}, fmt.Errorf("API config not found or not accessible") return ProviderResolution{}, fmt.Errorf("API config not found or not accessible")
} }
@@ -109,8 +77,8 @@ func ResolveProviderConfig(
// Resolve model: explicit > config default // Resolve model: explicit > config default
model := modelID model := modelID
if model == "" && modelDefault != nil { if model == "" && cfg.ModelDefault != "" {
model = *modelDefault model = cfg.ModelDefault
} }
if model == "" { if model == "" {
return ProviderResolution{}, fmt.Errorf("no model specified and no default model in config") return ProviderResolution{}, fmt.Errorf("no model specified and no default model in config")
@@ -118,10 +86,10 @@ func ResolveProviderConfig(
// Decrypt API key using the appropriate tier // Decrypt API key using the appropriate tier
key := "" key := ""
if len(apiKeyEnc) > 0 { if len(cfg.APIKeyEnc) > 0 {
if vault != nil { if vault != nil {
var err error var err error
key, err = vault.Decrypt(apiKeyEnc, keyNonce, keyScope, userID) key, err = vault.Decrypt(cfg.APIKeyEnc, cfg.KeyNonce, cfg.KeyScope, userID)
if err != nil { if err != nil {
if err == crypto.ErrVaultLocked { if err == crypto.ErrVaultLocked {
return ProviderResolution{}, fmt.Errorf("personal vault is locked — please log in again") return ProviderResolution{}, fmt.Errorf("personal vault is locked — please log in again")
@@ -130,40 +98,47 @@ func ResolveProviderConfig(
} }
} else { } else {
// No vault — key stored as raw bytes (unencrypted fallback) // No vault — key stored as raw bytes (unencrypted fallback)
key = string(apiKeyEnc) key = string(cfg.APIKeyEnc)
} }
} }
// Parse custom headers // Parse custom headers
customHeaders := make(map[string]string) customHeaders := make(map[string]string)
if customHeadersJSON != nil { if cfg.Headers != nil {
_ = json.Unmarshal(customHeadersJSON, &customHeaders) b, _ := json.Marshal(cfg.Headers)
_ = json.Unmarshal(b, &customHeaders)
} }
// Parse provider-specific settings // Parse provider-specific settings
providerSettings := make(map[string]interface{}) providerSettings := make(map[string]interface{})
if providerSettingsJSON != nil { if cfg.Settings != nil {
_ = json.Unmarshal(providerSettingsJSON, &providerSettings) for k, v := range cfg.Settings {
providerSettings[k] = v
}
} }
proxyMode := cfg.ProxyMode
if proxyMode == "" {
proxyMode = "system"
}
proxyURLStr := "" proxyURLStr := ""
if proxyURL != nil { if cfg.ProxyURL != nil {
proxyURLStr = *proxyURL proxyURLStr = *cfg.ProxyURL
} }
return ProviderResolution{ return ProviderResolution{
Config: providers.ProviderConfig{ Config: providers.ProviderConfig{
Endpoint: endpoint, Endpoint: cfg.Endpoint,
APIKey: key, APIKey: key,
CustomHeaders: customHeaders, CustomHeaders: customHeaders,
Settings: providerSettings, Settings: providerSettings,
ProxyMode: proxyMode, ProxyMode: proxyMode,
ProxyURL: proxyURLStr, ProxyURL: proxyURLStr,
}, },
ProviderID: providerID, ProviderID: cfg.Provider,
Model: model, Model: model,
ConfigID: configID, ConfigID: cfg.ID,
ProviderScope: providerScope, ProviderScope: cfg.Scope,
}, nil }, nil
} }

View File

@@ -59,7 +59,7 @@ func TestRouteRegistration(t *testing.T) {
protected.Use(middleware.Auth(cfg, stores.Users, userCache)) protected.Use(middleware.Auth(cfg, stores.Users, userCache))
// Channels (the route group that conflicts with workflow) // Channels (the route group that conflicts with workflow)
channels := NewChannelHandler() channels := NewChannelHandler(stores)
protected.GET("/channels", channels.ListChannels) protected.GET("/channels", channels.ListChannels)
protected.POST("/channels", channels.CreateChannel) protected.POST("/channels", channels.CreateChannel)
protected.GET("/channels/:id", channels.GetChannel) protected.GET("/channels/:id", channels.GetChannel)
@@ -93,7 +93,7 @@ func TestRouteRegistration(t *testing.T) {
protected.POST("/channels/:id/workflow/reject", wfInstH.Reject) protected.POST("/channels/:id/workflow/reject", wfInstH.Reject)
// Workflow assignments (v0.26.4) // Workflow assignments (v0.26.4)
wfAssignH := NewWorkflowAssignmentHandler() wfAssignH := NewWorkflowAssignmentHandler(stores)
protected.GET("/workflow-assignments/mine", wfAssignH.ListMine) protected.GET("/workflow-assignments/mine", wfAssignH.ListMine)
protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim) protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim)
protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete) protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete)

View File

@@ -1,7 +1,11 @@
package handlers package handlers
// settings.go — User profile, preferences, and password management.
//
// v0.29.0: Raw SQL replaced with UserStore methods.
import ( import (
"database/sql" "context"
"encoding/json" "encoding/json"
"log" "log"
"net/http" "net/http"
@@ -12,7 +16,7 @@ import (
"golang.org/x/crypto/bcrypt" "golang.org/x/crypto/bcrypt"
"git.gobha.me/xcaliber/chat-switchboard/crypto" "git.gobha.me/xcaliber/chat-switchboard/crypto"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
// ── Request Types ─────────────────────────── // ── Request Types ───────────────────────────
@@ -41,12 +45,13 @@ type profileResponse struct {
// SettingsHandler manages user profile and preferences. // SettingsHandler manages user profile and preferences.
type SettingsHandler struct { type SettingsHandler struct {
stores store.Stores
uekCache *crypto.UEKCache uekCache *crypto.UEKCache
} }
// NewSettingsHandler creates a new handler. // NewSettingsHandler creates a new handler.
func NewSettingsHandler(uekCache *crypto.UEKCache) *SettingsHandler { func NewSettingsHandler(stores store.Stores, uekCache *crypto.UEKCache) *SettingsHandler {
return &SettingsHandler{uekCache: uekCache} return &SettingsHandler{stores: stores, uekCache: uekCache}
} }
// ── Get Profile ───────────────────────────── // ── Get Profile ─────────────────────────────
@@ -54,31 +59,37 @@ func NewSettingsHandler(uekCache *crypto.UEKCache) *SettingsHandler {
func (h *SettingsHandler) GetProfile(c *gin.Context) { func (h *SettingsHandler) GetProfile(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
var p profileResponse user, err := h.stores.Users.GetByID(c.Request.Context(), userID)
var settingsRaw string if err != nil {
err := database.DB.QueryRow(database.Q(`
SELECT id, username, email, display_name, role, avatar_url,
COALESCE(NULLIF(settings::text, 'null'), '{}'),
created_at, last_login_at
FROM users WHERE id = $1
`), userID).Scan(
&p.ID, &p.Username, &p.Email, &p.DisplayName, &p.Role,
&p.Avatar, &settingsRaw,
database.ST(&p.CreatedAt), database.SNT(&p.LastLoginAt),
)
if err == sql.ErrNoRows {
c.JSON(http.StatusNotFound, gin.H{"error": "user not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "user not found"})
return return
} }
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load profile"}) var dn *string
return if user.DisplayName != "" {
dn = &user.DisplayName
}
var avatar *string
if user.AvatarURL != "" {
avatar = &user.AvatarURL
} }
p.Settings = make(map[string]interface{}) settings := make(map[string]interface{})
_ = json.Unmarshal([]byte(settingsRaw), &p.Settings) if user.Settings != nil {
settings = user.Settings
}
c.JSON(http.StatusOK, p) c.JSON(http.StatusOK, profileResponse{
ID: user.ID,
Username: user.Username,
Email: user.Email,
DisplayName: dn,
Role: user.Role,
Avatar: avatar,
Settings: settings,
CreatedAt: user.CreatedAt,
LastLoginAt: user.LastLoginAt,
})
} }
// ── Update Profile ────────────────────────── // ── Update Profile ──────────────────────────
@@ -93,11 +104,9 @@ func (h *SettingsHandler) UpdateProfile(c *gin.Context) {
} }
if req.DisplayName != nil { if req.DisplayName != nil {
_, err := database.DB.Exec( if err := h.stores.Users.Update(c.Request.Context(), userID, map[string]interface{}{
database.Q(`UPDATE users SET display_name = $1, updated_at = NOW() WHERE id = $2`), "display_name": *req.DisplayName,
*req.DisplayName, userID, }); err != nil {
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update display name"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update display name"})
return return
} }
@@ -105,12 +114,11 @@ func (h *SettingsHandler) UpdateProfile(c *gin.Context) {
if req.Email != nil { if req.Email != nil {
email := strings.ToLower(strings.TrimSpace(*req.Email)) email := strings.ToLower(strings.TrimSpace(*req.Email))
_, err := database.DB.Exec( err := h.stores.Users.Update(c.Request.Context(), userID, map[string]interface{}{
database.Q(`UPDATE users SET email = $1, updated_at = NOW() WHERE id = $2`), "email": email,
email, userID, })
)
if err != nil { if err != nil {
if database.IsUniqueViolation(err) { if isDuplicateErr(err) {
c.JSON(http.StatusConflict, gin.H{"error": "email already taken"}) c.JSON(http.StatusConflict, gin.H{"error": "email already taken"})
return return
} }
@@ -133,40 +141,31 @@ func (h *SettingsHandler) ChangePassword(c *gin.Context) {
return return
} }
// Verify current password user, err := h.stores.Users.GetByID(c.Request.Context(), userID)
var hash string
err := database.DB.QueryRow(
database.Q(`SELECT password_hash FROM users WHERE id = $1`), userID,
).Scan(&hash)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify password"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to verify password"})
return return
} }
if err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(req.CurrentPassword)); err != nil { if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.CurrentPassword)); err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "current password is incorrect"}) c.JSON(http.StatusUnauthorized, gin.H{"error": "current password is incorrect"})
return return
} }
// Hash new password
newHash, err := bcrypt.GenerateFromPassword([]byte(req.NewPassword), bcryptCost) newHash, err := bcrypt.GenerateFromPassword([]byte(req.NewPassword), bcryptCost)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to hash password"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to hash password"})
return return
} }
_, err = database.DB.Exec( if err := h.stores.Users.Update(c.Request.Context(), userID, map[string]interface{}{
database.Q(`UPDATE users SET password_hash = $1, updated_at = NOW() WHERE id = $2`), "password_hash": string(newHash),
string(newHash), userID, }); err != nil {
)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update password"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update password"})
return return
} }
// Re-wrap UEK with new password so personal BYOK keys remain accessible
h.rewrapVault(userID, req.CurrentPassword, req.NewPassword) h.rewrapVault(userID, req.CurrentPassword, req.NewPassword)
c.JSON(http.StatusOK, gin.H{"message": "password updated"}) c.JSON(http.StatusOK, gin.H{"message": "password updated"})
} }
@@ -175,20 +174,16 @@ func (h *SettingsHandler) ChangePassword(c *gin.Context) {
func (h *SettingsHandler) GetSettings(c *gin.Context) { func (h *SettingsHandler) GetSettings(c *gin.Context) {
userID := getUserID(c) userID := getUserID(c)
var settingsRaw string user, err := h.stores.Users.GetByID(c.Request.Context(), userID)
// 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 COALESCE(NULLIF(settings::text, 'null'), '{}') FROM users WHERE id = $1`), userID,
).Scan(&settingsRaw)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load settings"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load settings"})
return return
} }
settings := make(map[string]interface{}) settings := make(map[string]interface{})
_ = json.Unmarshal([]byte(settingsRaw), &settings) if user.Settings != nil {
settings = user.Settings
}
c.JSON(http.StatusOK, gin.H{"settings": settings}) c.JSON(http.StatusOK, gin.H{"settings": settings})
} }
@@ -210,28 +205,7 @@ func (h *SettingsHandler) UpdateSettings(c *gin.Context) {
return return
} }
// JSONB merge — existing keys preserved, incoming keys overwrite. if err := h.stores.Users.MergeSettings(c.Request.Context(), userID, patch); err != nil {
// 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(
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 = (
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 {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update settings"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update settings"})
return return
} }
@@ -241,25 +215,16 @@ func (h *SettingsHandler) UpdateSettings(c *gin.Context) {
// ── Vault Re-wrap ────────────────────────── // ── Vault Re-wrap ──────────────────────────
// rewrapVault decrypts the UEK with the old password and re-encrypts it with
// the new password. This keeps personal BYOK keys accessible after a password
// change. Uses a new salt for forward secrecy.
func (h *SettingsHandler) rewrapVault(userID, oldPassword, newPassword string) { func (h *SettingsHandler) rewrapVault(userID, oldPassword, newPassword string) {
if h.uekCache == nil { if h.uekCache == nil {
return return
} }
var vaultSet bool vaultSet, encUEK, salt, nonce, err := h.stores.Users.GetVaultKeys(c_bg(), userID)
var encUEK, salt, nonce []byte
err := database.DB.QueryRow(database.Q(`
SELECT vault_set, encrypted_uek, uek_salt, uek_nonce
FROM users WHERE id = $1
`), userID).Scan(&vaultSet, &encUEK, &salt, &nonce)
if err != nil || !vaultSet || len(encUEK) == 0 { if err != nil || !vaultSet || len(encUEK) == 0 {
return // No vault to re-wrap return
} }
// Decrypt UEK with old password
oldPDK := crypto.DeriveKeyFromPassword(oldPassword, salt) oldPDK := crypto.DeriveKeyFromPassword(oldPassword, salt)
uek, err := crypto.UnwrapUEK(encUEK, nonce, oldPDK) uek, err := crypto.UnwrapUEK(encUEK, nonce, oldPDK)
if err != nil { if err != nil {
@@ -267,7 +232,6 @@ func (h *SettingsHandler) rewrapVault(userID, oldPassword, newPassword string) {
return return
} }
// Re-wrap with new password using fresh salt
newSalt, err := crypto.GenerateSalt() newSalt, err := crypto.GenerateSalt()
if err != nil { if err != nil {
log.Printf("⚠ Vault re-wrap failed for user %s (salt): %v", userID, err) log.Printf("⚠ Vault re-wrap failed for user %s (salt): %v", userID, err)
@@ -281,17 +245,15 @@ func (h *SettingsHandler) rewrapVault(userID, oldPassword, newPassword string) {
return return
} }
_, err = database.DB.Exec(database.Q(` if err := h.stores.Users.UpdateVaultKeys(c_bg(), userID, newEncUEK, newSalt, newNonce); err != nil {
UPDATE users
SET encrypted_uek = $1, uek_salt = $2, uek_nonce = $3, updated_at = NOW()
WHERE id = $4
`), newEncUEK, newSalt, newNonce, userID)
if err != nil {
log.Printf("⚠ Vault re-wrap failed for user %s (persist): %v", userID, err) log.Printf("⚠ Vault re-wrap failed for user %s (persist): %v", userID, err)
return return
} }
// Session cache: UEK itself hasn't changed, just its wrapper
h.uekCache.Store(userID, uek) h.uekCache.Store(userID, uek)
log.Printf(" 🔐 Vault re-wrapped for user %s", userID) log.Printf(" 🔐 Vault re-wrapped for user %s", userID)
} }
// c_bg returns context.Background() — short alias for vault operations
// that run outside the request lifecycle.
func c_bg() context.Context { return context.Background() }

View File

@@ -1,45 +1,44 @@
package handlers package handlers
// summarize.go — Manual conversation summarization via HTTP.
//
// v0.29.0: Raw SQL replaced with ChannelStore methods.
import ( import (
"net/http" "net/http"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/compaction" "git.gobha.me/xcaliber/chat-switchboard/compaction"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/roles" "git.gobha.me/xcaliber/chat-switchboard/roles"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
// SummarizeHandler handles manual conversation summarization via HTTP. // SummarizeHandler handles manual conversation summarization via HTTP.
type SummarizeHandler struct { type SummarizeHandler struct {
stores store.Stores
compaction *compaction.Service compaction *compaction.Service
} }
// NewSummarizeHandler creates a new handler backed by the compaction service. // NewSummarizeHandler creates a new handler backed by the compaction service.
func NewSummarizeHandler(svc *compaction.Service) *SummarizeHandler { func NewSummarizeHandler(stores store.Stores, svc *compaction.Service) *SummarizeHandler {
return &SummarizeHandler{compaction: svc} return &SummarizeHandler{stores: stores, compaction: svc}
} }
// ── Summarize & Continue ────────────────── // ── Summarize & Continue ──────────────────
// POST /channels/:id/summarize // POST /channels/:id/summarize
//
// User-triggered compaction: calls the utility role to summarize the
// conversation history, inserts the summary as a tree node, and returns it.
func (h *SummarizeHandler) Summarize(c *gin.Context) { func (h *SummarizeHandler) Summarize(c *gin.Context) {
channelID := c.Param("id") channelID := c.Param("id")
userID := getUserID(c) userID := getUserID(c)
// ── Verify channel ownership ── // ── Verify channel ownership ──
var ownerID string ch, err := h.stores.Channels.GetByID(c.Request.Context(), channelID)
err := database.DB.QueryRow( if err != nil || ch == nil {
database.Q(`SELECT user_id FROM channels WHERE id = $1`), channelID,
).Scan(&ownerID)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "channel not found"})
return return
} }
if ownerID != userID { if ch.UserID != userID {
c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"}) c.JSON(http.StatusForbidden, gin.H{"error": "not your channel"})
return return
} }
@@ -73,7 +72,6 @@ func (h *SummarizeHandler) Summarize(c *gin.Context) {
Trigger: "manual", Trigger: "manual",
}) })
if err != nil { if err != nil {
// Map specific errors to HTTP status codes
switch err.Error() { switch err.Error() {
case "conversation too short to summarize": case "conversation too short to summarize":
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})

View File

@@ -689,11 +689,11 @@ func TestTask_WorkflowTypeRejected(t *testing.T) {
// Task Type RBAC (v0.28.7) // Task Type RBAC (v0.28.7)
// ═══════════════════════════════════════════════ // ═══════════════════════════════════════════════
func TestTask_StarlarkTypeRejected(t *testing.T) { func TestTask_StarlarkTypeRequiresPackageID(t *testing.T) {
h := setupHarness(t) h := setupHarness(t)
_, token := h.createAdminUser("staruser", "staruser@test.com") _, token := h.createAdminUser("staruser", "staruser@test.com")
// Even admin cannot create starlark tasks — executor doesn't exist yet // Starlark tasks require system_function (package_id)
w := h.request("POST", "/api/v1/tasks", token, map[string]interface{}{ w := h.request("POST", "/api/v1/tasks", token, map[string]interface{}{
"name": "Starlark Task", "name": "Starlark Task",
"task_type": "starlark", "task_type": "starlark",
@@ -702,11 +702,11 @@ func TestTask_StarlarkTypeRejected(t *testing.T) {
"model_id": "m", "model_id": "m",
}) })
if w.Code != http.StatusBadRequest { if w.Code != http.StatusBadRequest {
t.Fatalf("starlark task_type: want 400, got %d: %s", w.Code, w.Body.String()) t.Fatalf("starlark task_type without package_id: want 400, got %d: %s", w.Code, w.Body.String())
} }
body := w.Body.String() body := w.Body.String()
if !strings.Contains(body, "v0.29.0") { if !strings.Contains(body, "system_function") && !strings.Contains(body, "package_id") {
t.Fatalf("expected 'v0.29.0' in error, got: %s", body) t.Fatalf("expected 'system_function' or 'package_id' in error, got: %s", body)
} }
} }

View File

@@ -133,12 +133,30 @@ func (h *TaskHandler) Create(c *gin.Context) {
} }
} }
// v0.28.7: Starlark tasks — pre-positioned gate for v0.29.0. // v0.29.0: Starlark tasks — run sandboxed extension scripts.
// The Starlark executor does not exist yet. Reject all creation with // Requires task.starlark permission. system_function holds package ID.
// a clear message rather than a cryptic CHECK constraint failure.
if t.TaskType == "starlark" { if t.TaskType == "starlark" {
c.JSON(http.StatusBadRequest, gin.H{"error": "starlark task execution requires v0.29.0 — not yet available"}) if c.GetString("role") != "admin" {
return perms := middleware.GetResolvedPermissions(c)
if perms == nil || !perms[auth.PermTaskStarlark] {
c.JSON(http.StatusForbidden, gin.H{"error": "permission required: " + auth.PermTaskStarlark})
return
}
}
if t.SystemFunction == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "system_function (package_id) is required for starlark tasks"})
return
}
// Verify package exists and is starlark tier
pkg, err := h.stores.Packages.Get(c.Request.Context(), t.SystemFunction)
if err != nil || pkg == nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "package not found: " + t.SystemFunction})
return
}
if pkg.Tier != models.ExtTierStarlark {
c.JSON(http.StatusBadRequest, gin.H{"error": "package " + t.SystemFunction + " is not a starlark package"})
return
}
} }
// v0.28.0-audit: Workflow task execution is not yet implemented. // v0.28.0-audit: Workflow task execution is not yet implemented.

View File

@@ -1,6 +1,7 @@
package handlers package handlers
import ( import (
"context"
"encoding/json" "encoding/json"
"log" "log"
"net/http" "net/http"
@@ -8,9 +9,9 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities" capspkg "git.gobha.me/xcaliber/chat-switchboard/capabilities"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/providers" "git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
// ── Team Provider Handlers ────────────────── // ── Team Provider Handlers ──────────────────
@@ -18,19 +19,13 @@ import (
// ListTeamProviders returns API configs scoped to a team. // ListTeamProviders returns API configs scoped to a team.
func (h *TeamHandler) ListTeamProviders(c *gin.Context) { func (h *TeamHandler) ListTeamProviders(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
ctx := c.Request.Context()
rows, err := database.DB.Query(database.Q(` configs, err := h.stores.Providers.ListAllForTeam(ctx, teamID)
SELECT id, name, provider, endpoint, api_key_enc,
model_default, config, is_active, is_private, created_at, updated_at
FROM provider_configs
WHERE scope = 'team' AND owner_id = $1
ORDER BY name ASC
`), teamID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list team providers"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list team providers"})
return return
} }
defer rows.Close()
type teamProvider struct { type teamProvider struct {
ID string `json:"id"` ID string `json:"id"`
@@ -46,23 +41,34 @@ func (h *TeamHandler) ListTeamProviders(c *gin.Context) {
UpdatedAt string `json:"updated_at"` UpdatedAt string `json:"updated_at"`
} }
configs := make([]teamProvider, 0) result := make([]teamProvider, 0, len(configs))
for rows.Next() { for _, cfg := range configs {
var p teamProvider var md *string
var apiKeyEnc []byte if cfg.ModelDefault != "" {
var configRaw string md = &cfg.ModelDefault
if err := rows.Scan(&p.ID, &p.Name, &p.Provider, &p.Endpoint, &apiKeyEnc,
&p.ModelDefault, &configRaw, &p.IsActive, &p.IsPrivate, &p.CreatedAt, &p.UpdatedAt); err != nil {
continue
} }
p.HasKey = len(apiKeyEnc) > 0 cfgMap := map[string]interface{}{}
p.Config = parseJSONBConfig(configRaw) if cfg.Config != nil {
configs = append(configs, p) cfgMap = cfg.Config
}
result = append(result, teamProvider{
ID: cfg.ID,
Name: cfg.Name,
Provider: cfg.Provider,
Endpoint: cfg.Endpoint,
HasKey: cfg.HasKey(),
ModelDefault: md,
Config: cfgMap,
IsActive: cfg.IsActive,
IsPrivate: cfg.IsPrivate,
CreatedAt: cfg.CreatedAt.Format("2006-01-02T15:04:05Z"),
UpdatedAt: cfg.UpdatedAt.Format("2006-01-02T15:04:05Z"),
})
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
"data": configs, "data": result,
"allow_team_providers": isTeamProvidersAllowed(teamID), "allow_team_providers": isTeamProvidersAllowed(h.stores, teamID),
}) })
} }
@@ -70,7 +76,7 @@ func (h *TeamHandler) ListTeamProviders(c *gin.Context) {
func (h *TeamHandler) CreateTeamProvider(c *gin.Context) { func (h *TeamHandler) CreateTeamProvider(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
if !isTeamProvidersAllowed(teamID) { if !isTeamProvidersAllowed(h.stores, teamID) {
c.JSON(http.StatusForbidden, gin.H{"error": "team providers are not enabled for this team"}) c.JSON(http.StatusForbidden, gin.H{"error": "team providers are not enabled for this team"})
return return
} }
@@ -98,18 +104,6 @@ func (h *TeamHandler) CreateTeamProvider(c *gin.Context) {
return return
} }
configJSON := "{}"
if req.Config != nil {
b, _ := json.Marshal(req.Config)
configJSON = string(b)
}
headersJSON := "{}"
if req.Headers != nil {
b, _ := json.Marshal(req.Headers)
headersJSON = string(b)
}
// Encrypt the API key for team scope // Encrypt the API key for team scope
var apiKeyEnc, keyNonce []byte var apiKeyEnc, keyNonce []byte
if req.APIKey != "" { if req.APIKey != "" {
@@ -125,27 +119,43 @@ func (h *TeamHandler) CreateTeamProvider(c *gin.Context) {
} }
} }
id, err := database.InsertReturningID(` headersMap := models.JSONMap{}
INSERT INTO provider_configs (scope, owner_id, name, provider, endpoint, if req.Headers != nil {
api_key_enc, key_nonce, key_scope, model_default, config, is_private, headers) for k, v := range req.Headers {
VALUES ($1, $2, $3, $4, $5, $6, $7, 'team', $8, $9::jsonb, $10, $11::jsonb) headersMap[k] = v
RETURNING id }
`, "team", teamID, req.Name, req.Provider, req.Endpoint, }
apiKeyEnc, keyNonce, req.ModelDefault, configJSON, req.IsPrivate, headersJSON,
) cfg := &models.ProviderConfig{
if err != nil { Scope: models.ScopeTeam,
OwnerID: &teamID,
Name: req.Name,
Provider: req.Provider,
Endpoint: req.Endpoint,
APIKeyEnc: apiKeyEnc,
KeyNonce: keyNonce,
KeyScope: "team",
ModelDefault: req.ModelDefault,
Config: models.JSONMap(req.Config),
Headers: headersMap,
IsActive: true,
IsPrivate: req.IsPrivate,
}
if err := h.stores.Providers.Create(c.Request.Context(), cfg); err != nil {
log.Printf("[WARN] Failed to create team provider: %v", err) log.Printf("[WARN] Failed to create team provider: %v", err)
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create provider"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create provider"})
return return
} }
c.JSON(http.StatusCreated, gin.H{"id": id}) c.JSON(http.StatusCreated, gin.H{"id": cfg.ID})
} }
// UpdateTeamProvider updates a team-scoped API config. // UpdateTeamProvider updates a team-scoped API config.
func (h *TeamHandler) UpdateTeamProvider(c *gin.Context) { func (h *TeamHandler) UpdateTeamProvider(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
providerID := c.Param("id") providerID := c.Param("id")
ctx := c.Request.Context()
var req struct { var req struct {
Name *string `json:"name,omitempty"` Name *string `json:"name,omitempty"`
@@ -163,29 +173,23 @@ func (h *TeamHandler) UpdateTeamProvider(c *gin.Context) {
} }
// Verify provider belongs to this team // Verify provider belongs to this team
var count int existing, err := h.stores.Providers.GetByID(ctx, providerID)
database.DB.QueryRow(database.Q(`SELECT COUNT(*) FROM provider_configs WHERE id = $1 AND scope = 'team' AND owner_id = $2`), providerID, teamID).Scan(&count) if err != nil || existing.Scope != models.ScopeTeam || (existing.OwnerID != nil && *existing.OwnerID != teamID) {
if count == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "provider not found in this team"}) c.JSON(http.StatusNotFound, gin.H{"error": "provider not found in this team"})
return return
} }
// Build dynamic update using ? placeholders (works on both dialects) // Build patch
setClauses := []string{"updated_at = " + database.Q("NOW()")} patch := models.ProviderConfigPatch{}
args := []interface{}{}
fieldCount := 0 fieldCount := 0
addSet := func(col string, val interface{}) { if req.Name != nil {
setClauses = append(setClauses, col+" = ?") patch.Name = req.Name
args = append(args, val)
fieldCount++ fieldCount++
} }
if req.Name != nil {
addSet("name", *req.Name)
}
if req.Endpoint != nil { if req.Endpoint != nil {
addSet("endpoint", *req.Endpoint) patch.Endpoint = req.Endpoint
fieldCount++
} }
if req.APIKey != nil && *req.APIKey != "" { if req.APIKey != nil && *req.APIKey != "" {
if h.vault != nil { if h.vault != nil {
@@ -194,28 +198,36 @@ func (h *TeamHandler) UpdateTeamProvider(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to encrypt API key"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to encrypt API key"})
return return
} }
addSet("api_key_enc", enc) patch.APIKeyEnc = enc
addSet("key_nonce", nonce) patch.KeyNonce = nonce
} else { } else {
addSet("api_key_enc", []byte(*req.APIKey)) patch.APIKeyEnc = []byte(*req.APIKey)
} }
fieldCount++
} }
if req.ModelDefault != nil { if req.ModelDefault != nil {
addSet("model_default", *req.ModelDefault) patch.ModelDefault = req.ModelDefault
fieldCount++
} }
if req.IsActive != nil { if req.IsActive != nil {
addSet("is_active", *req.IsActive) patch.IsActive = req.IsActive
fieldCount++
} }
if req.IsPrivate != nil { if req.IsPrivate != nil {
addSet("is_private", *req.IsPrivate) patch.IsPrivate = req.IsPrivate
fieldCount++
} }
if req.Config != nil { if req.Config != nil {
b, _ := json.Marshal(req.Config) patch.Config = models.JSONMap(req.Config)
addSet("config", string(b)) fieldCount++
} }
if req.Headers != nil { if req.Headers != nil {
b, _ := json.Marshal(req.Headers) headersMap := models.JSONMap{}
addSet("headers", string(b)) for k, v := range req.Headers {
headersMap[k] = v
}
patch.Headers = headersMap
fieldCount++
} }
if fieldCount == 0 { if fieldCount == 0 {
@@ -223,24 +235,7 @@ func (h *TeamHandler) UpdateTeamProvider(c *gin.Context) {
return return
} }
args = append(args, providerID, teamID) if err := h.stores.Providers.Update(ctx, providerID, patch); err != nil {
// Build final query with ? placeholders, then convert to $N for Postgres
query := "UPDATE provider_configs SET "
for i, s := range setClauses {
if i > 0 {
query += ", "
}
query += s
}
query += " WHERE id = ? AND scope = 'team' AND owner_id = ?"
if database.IsPostgres() {
query = convertPlaceholders(query)
}
_, err := database.DB.Exec(query, args...)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update provider"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to update provider"})
return return
} }
@@ -253,16 +248,12 @@ func (h *TeamHandler) DeleteTeamProvider(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
providerID := c.Param("id") providerID := c.Param("id")
result, err := database.DB.Exec(database.Q(` n, err := h.stores.Providers.DeleteByIDAndTeam(c.Request.Context(), providerID, teamID)
DELETE FROM provider_configs WHERE id = $1 AND scope = 'team' AND owner_id = $2
`), providerID, teamID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete provider"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to delete provider"})
return return
} }
if n == 0 {
rows, _ := result.RowsAffected()
if rows == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "provider not found in this team"}) c.JSON(http.StatusNotFound, gin.H{"error": "provider not found in this team"})
return return
} }
@@ -274,46 +265,46 @@ func (h *TeamHandler) DeleteTeamProvider(c *gin.Context) {
func (h *TeamHandler) ListTeamProviderModels(c *gin.Context) { func (h *TeamHandler) ListTeamProviderModels(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
providerID := c.Param("id") providerID := c.Param("id")
ctx := c.Request.Context()
var name, providerType, endpoint string cfg, err := h.stores.Providers.GetByID(ctx, providerID)
var apiKeyEnc, keyNonce []byte if err != nil || cfg.Scope != models.ScopeTeam || !cfg.IsActive {
var keyScope string c.JSON(http.StatusNotFound, gin.H{"error": "provider not found"})
var headersJSON []byte return
err := database.DB.QueryRow(database.Q(` }
SELECT name, provider, endpoint, api_key_enc, key_nonce, key_scope, headers if cfg.OwnerID == nil || *cfg.OwnerID != teamID {
FROM provider_configs
WHERE id = $1 AND scope = 'team' AND owner_id = $2 AND is_active = true
`), providerID, teamID).Scan(&name, &providerType, &endpoint, &apiKeyEnc, &keyNonce, &keyScope, &headersJSON)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "provider not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "provider not found"})
return return
} }
provider, err := providers.Get(providerType) provider, err := providers.Get(cfg.Provider)
if err != nil { if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported provider"}) c.JSON(http.StatusBadRequest, gin.H{"error": "unsupported provider"})
return return
} }
key := "" key := ""
if len(apiKeyEnc) > 0 { if cfg.HasKey() {
if h.vault != nil { if h.vault != nil {
var err error var err error
key, err = h.vault.Decrypt(apiKeyEnc, keyNonce, keyScope, "") key, err = h.vault.Decrypt(cfg.APIKeyEnc, cfg.KeyNonce, cfg.KeyScope, "")
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to decrypt API key"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to decrypt API key"})
return return
} }
} else { } else {
key = string(apiKeyEnc) key = string(cfg.APIKeyEnc)
} }
} }
var customHeaders map[string]string var customHeaders map[string]string
_ = json.Unmarshal(headersJSON, &customHeaders) if cfg.Headers != nil {
b, _ := json.Marshal(cfg.Headers)
_ = json.Unmarshal(b, &customHeaders)
}
modelList, err := provider.ListModels(c.Request.Context(), providers.ProviderConfig{ modelList, err := provider.ListModels(ctx, providers.ProviderConfig{
Endpoint: endpoint, Endpoint: cfg.Endpoint,
APIKey: key, APIKey: key,
CustomHeaders: customHeaders, CustomHeaders: customHeaders,
}) })
@@ -335,7 +326,7 @@ func (h *TeamHandler) ListTeamProviderModels(c *gin.Context) {
out = append(out, modelInfo{ID: m.ID, Type: m.Type, Capabilities: caps}) out = append(out, modelInfo{ID: m.ID, Type: m.Type, Capabilities: caps})
} }
c.JSON(http.StatusOK, gin.H{"models": out, "provider": name}) c.JSON(http.StatusOK, gin.H{"models": out, "provider": cfg.Name})
} }
// parseJSONBConfig parses a JSONB text string into a map. // parseJSONBConfig parses a JSONB text string into a map.
@@ -351,32 +342,30 @@ func parseJSONBConfig(raw string) map[string]interface{} {
} }
// isTeamProvidersAllowed checks if team providers are enabled. // isTeamProvidersAllowed checks if team providers are enabled.
func isTeamProvidersAllowed(teamID string) bool { func isTeamProvidersAllowed(stores store.Stores, teamID string) bool {
if database.DB == nil { if stores.GlobalConfig == nil {
return false return false
} }
var globalVal string ctx := context.Background()
err := database.DB.QueryRow(`
SELECT value FROM global_settings WHERE key = 'allow_team_providers' // Check global setting
`).Scan(&globalVal) globalVal, err := stores.GlobalConfig.GetString(ctx, "allow_team_providers")
if err == nil && globalVal == "false" { if err == nil && globalVal == "false" {
return false return false
} }
var settingsJSON []byte // Check team-level setting
err = database.DB.QueryRow(database.Q(`SELECT settings FROM teams WHERE id = $1`), teamID).Scan(&settingsJSON) team, err := stores.Teams.GetByID(ctx, teamID)
if err != nil { if err != nil {
return true return true // fail open
} }
var settings map[string]interface{} if team.Settings != nil {
if err := json.Unmarshal(settingsJSON, &settings); err != nil { if v, ok := team.Settings["allow_team_providers"]; ok {
return true if bVal, ok := v.(bool); ok {
} return bVal
if v, ok := settings["allow_team_providers"]; ok { }
if b, ok := v.(bool); ok {
return b
} }
} }

View File

@@ -1,14 +1,13 @@
package handlers package handlers
import ( import (
"context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"fmt" "fmt"
"log" "log"
"net/http" "net/http"
"strconv" "strconv"
"strings"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -152,68 +151,44 @@ func (h *TeamHandler) UpdateTeam(c *gin.Context) {
return return
} }
// Build dynamic update // Build fields map for store.Update (simple scalar fields)
sets := []string{} fields := map[string]interface{}{}
args := []interface{}{}
argN := 1
addArg := func(col string, val interface{}) {
if database.IsSQLite() {
sets = append(sets, col+" = ?")
} else {
sets = append(sets, col+" = $"+strconv.Itoa(argN))
}
args = append(args, val)
argN++
}
if req.Name != nil { if req.Name != nil {
addArg("name", *req.Name) fields["name"] = *req.Name
} }
if req.Description != nil { if req.Description != nil {
addArg("description", *req.Description) fields["description"] = *req.Description
} }
if req.IsActive != nil { if req.IsActive != nil {
addArg("is_active", *req.IsActive) fields["is_active"] = *req.IsActive
}
if req.Settings != nil {
if database.IsSQLite() {
// SQLite: json_patch for merge
sets = append(sets, "settings = json_patch(COALESCE(settings, '{}'), ?)")
} else {
sets = append(sets, "settings = COALESCE(settings, '{}'::jsonb) || $"+strconv.Itoa(argN)+"::jsonb")
}
args = append(args, *req.Settings)
argN++
} }
if len(sets) == 0 { hasFields := len(fields) > 0
hasSettings := req.Settings != nil
if !hasFields && !hasSettings {
c.JSON(http.StatusBadRequest, gin.H{"error": "no fields to update"}) c.JSON(http.StatusBadRequest, gin.H{"error": "no fields to update"})
return return
} }
var whereClause string ctx := c.Request.Context()
if database.IsSQLite() {
whereClause = " WHERE id = ?"
} else {
whereClause = " WHERE id = $" + strconv.Itoa(argN)
}
args = append(args, teamID)
query := "UPDATE teams SET " + strings.Join(sets, ", ") + whereClause if hasFields {
if err := h.stores.Teams.Update(ctx, teamID, fields); err != nil {
res, err := database.DB.Exec(query, args...) if database.IsUniqueViolation(err) {
if err != nil { c.JSON(http.StatusConflict, gin.H{"error": "team name already exists"})
if database.IsUniqueViolation(err) { return
c.JSON(http.StatusConflict, gin.H{"error": "team name already exists"}) }
c.JSON(http.StatusInternalServerError, gin.H{"error": "update failed"})
return return
} }
c.JSON(http.StatusInternalServerError, gin.H{"error": "update failed"})
return
} }
if n, _ := res.RowsAffected(); n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "team not found"}) if hasSettings {
return if err := h.stores.Teams.MergeSettings(ctx, teamID, *req.Settings); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "settings update failed"})
return
}
} }
c.JSON(http.StatusOK, gin.H{"ok": true}) c.JSON(http.StatusOK, gin.H{"ok": true})
@@ -266,6 +241,7 @@ func (h *TeamHandler) ListMembers(c *gin.Context) {
func (h *TeamHandler) AddMember(c *gin.Context) { func (h *TeamHandler) AddMember(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
ctx := c.Request.Context()
var req addMemberRequest var req addMemberRequest
if err := c.ShouldBindJSON(&req); err != nil { if err := c.ShouldBindJSON(&req); err != nil {
@@ -274,25 +250,20 @@ func (h *TeamHandler) AddMember(c *gin.Context) {
} }
// Verify team exists // Verify team exists
var exists bool exists, _ := h.stores.Teams.Exists(ctx, teamID)
database.DB.QueryRow(database.Q(`SELECT EXISTS(SELECT 1 FROM teams WHERE id = $1)`), teamID).Scan(&exists)
if !exists { if !exists {
c.JSON(http.StatusNotFound, gin.H{"error": "team not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "team not found"})
return return
} }
// Verify user exists // Verify user exists
database.DB.QueryRow(database.Q(`SELECT EXISTS(SELECT 1 FROM users WHERE id = $1)`), req.UserID).Scan(&exists) userExists, _ := h.stores.Users.Exists(ctx, req.UserID)
if !exists { if !userExists {
c.JSON(http.StatusBadRequest, gin.H{"error": "user not found"}) c.JSON(http.StatusBadRequest, gin.H{"error": "user not found"})
return return
} }
id, err := database.InsertReturningID(` id, err := h.stores.Teams.AddMemberReturningID(ctx, teamID, req.UserID, req.Role)
INSERT INTO team_members (team_id, user_id, role)
VALUES ($1, $2, $3)
RETURNING id
`, teamID, req.UserID, req.Role)
if err != nil { if err != nil {
if database.IsUniqueViolation(err) { if database.IsUniqueViolation(err) {
c.JSON(http.StatusConflict, gin.H{"error": "user is already a member"}) c.JSON(http.StatusConflict, gin.H{"error": "user is already a member"})
@@ -320,14 +291,12 @@ func (h *TeamHandler) UpdateMember(c *gin.Context) {
return return
} }
res, err := database.DB.Exec(database.Q(` n, err := h.stores.Teams.UpdateMemberRoleByID(c.Request.Context(), memberID, teamID, req.Role)
UPDATE team_members SET role = $1 WHERE id = $2 AND team_id = $3
`), req.Role, memberID, teamID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "update failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "update failed"})
return return
} }
if n, _ := res.RowsAffected(); n == 0 { if n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "member not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "member not found"})
return return
} }
@@ -344,12 +313,12 @@ func (h *TeamHandler) RemoveMember(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
memberID := c.Param("memberId") memberID := c.Param("memberId")
res, err := database.DB.Exec(database.Q(`DELETE FROM team_members WHERE id = $1 AND team_id = $2`), memberID, teamID) n, err := h.stores.Teams.DeleteMemberByID(c.Request.Context(), memberID, teamID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "remove failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "remove failed"})
return return
} }
if n, _ := res.RowsAffected(); n == 0 { if n == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "member not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "member not found"})
return return
} }
@@ -384,6 +353,7 @@ func (h *TeamHandler) MyTeams(c *gin.Context) {
// GET /api/v1/teams/:teamId/models // GET /api/v1/teams/:teamId/models
func (h *TeamHandler) ListAvailableModels(c *gin.Context) { func (h *TeamHandler) ListAvailableModels(c *gin.Context) {
teamID := getTeamID(c) teamID := getTeamID(c)
ctx := c.Request.Context()
type availableModel struct { type availableModel struct {
ID string `json:"id"` ID string `json:"id"`
@@ -395,79 +365,62 @@ func (h *TeamHandler) ListAvailableModels(c *gin.Context) {
Source string `json:"source"` Source string `json:"source"`
} }
models := make([]availableModel, 0) result := make([]availableModel, 0)
// ── 1. Global admin models (synced in model_catalog) ── // ── 1. Global admin models (synced in model_catalog) ──
rows, err := database.DB.Query(database.Q(` catalogModels, err := h.stores.Catalog.ListTeamAvailable(ctx)
SELECT mc.id, mc.model_id, mc.display_name, mc.visibility,
ac.provider, ac.name as provider_name
FROM model_catalog mc
JOIN provider_configs ac ON mc.provider_config_id = ac.id
WHERE mc.visibility IN ('enabled', 'team')
AND ac.is_active = true AND ac.scope = 'global'
ORDER BY ac.name, mc.model_id
`))
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "query failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "query failed"})
return return
} }
defer rows.Close() for _, cm := range catalogModels {
result = append(result, availableModel{
for rows.Next() { ID: cm.ID,
var m availableModel ModelID: cm.ModelID,
if err := rows.Scan(&m.ID, &m.ModelID, &m.DisplayName, &m.Visibility, DisplayName: cm.DisplayName,
&m.Provider, &m.ProviderName); err != nil { Visibility: cm.Visibility,
continue Provider: cm.Provider,
} ProviderName: cm.ProviderName,
m.Source = "global" Source: "global",
models = append(models, m) })
} }
// ── 2. Team provider models (live query) ── // ── 2. Team provider models (live query) ──
teamRows, err := database.DB.Query(database.Q(` teamConfigs, err := h.stores.Providers.ListForTeam(ctx, teamID)
SELECT id, name, provider, endpoint, api_key_enc, headers
FROM provider_configs
WHERE scope = 'team' AND owner_id = $1 AND is_active = true
`), teamID)
if err == nil { if err == nil {
defer teamRows.Close() for _, cfg := range teamConfigs {
for teamRows.Next() { provider, pErr := providers.Get(cfg.Provider)
var cfgID, name, providerID, endpoint string
var apiKey *string
var headersJSON []byte
if err := teamRows.Scan(&cfgID, &name, &providerID, &endpoint, &apiKey, &headersJSON); err != nil {
continue
}
provider, pErr := providers.Get(providerID)
if pErr != nil { if pErr != nil {
continue continue
} }
key := "" key := ""
if apiKey != nil { if cfg.HasKey() {
key = *apiKey key = string(cfg.APIKeyEnc)
} }
var customHeaders map[string]string var customHeaders map[string]string
_ = json.Unmarshal(headersJSON, &customHeaders) if cfg.Headers != nil {
b, _ := json.Marshal(cfg.Headers)
_ = json.Unmarshal(b, &customHeaders)
}
provModels, lErr := provider.ListModels(c.Request.Context(), providers.ProviderConfig{ provModels, lErr := provider.ListModels(ctx, providers.ProviderConfig{
Endpoint: endpoint, Endpoint: cfg.Endpoint,
APIKey: key, APIKey: key,
CustomHeaders: customHeaders, CustomHeaders: customHeaders,
}) })
if lErr != nil { if lErr != nil {
log.Printf("[models] team provider %q list failed: %v", name, lErr) log.Printf("[models] team provider %q list failed: %v", cfg.Name, lErr)
continue continue
} }
for _, pm := range provModels { for _, pm := range provModels {
models = append(models, availableModel{ result = append(result, availableModel{
ID: pm.ID, ID: pm.ID,
ModelID: pm.ID, ModelID: pm.ID,
Provider: providerID, Provider: cfg.Provider,
ProviderName: name, ProviderName: cfg.Name,
Visibility: "enabled", Visibility: "enabled",
Source: "team", Source: "team",
}) })
@@ -475,7 +428,7 @@ func (h *TeamHandler) ListAvailableModels(c *gin.Context) {
} }
} }
c.JSON(http.StatusOK, gin.H{"models": models}) c.JSON(http.StatusOK, gin.H{"models": result})
} }
// ── Helpers ───────────────────────────────── // ── Helpers ─────────────────────────────────
@@ -488,68 +441,25 @@ func getTeamID(c *gin.Context) string {
return c.Param("id") return c.Param("id")
} }
// IsTeamAdmin checks if a user is an admin of the given team.
func IsTeamAdmin(userID, teamID string) bool {
var role string
err := database.DB.QueryRow(database.Q(`
SELECT role FROM team_members WHERE team_id = $1 AND user_id = $2
`), teamID, userID).Scan(&role)
return err == nil && role == "admin"
}
// IsTeamMember checks if a user belongs to the given team (any role).
func IsTeamMember(userID, teamID string) bool {
var exists bool
database.DB.QueryRow(database.Q(`
SELECT EXISTS(SELECT 1 FROM team_members WHERE team_id = $1 AND user_id = $2)
`), teamID, userID).Scan(&exists)
return exists
}
// enforcePrivateProviderPolicy checks if a user belongs to any team that // enforcePrivateProviderPolicy checks if a user belongs to any team that
// requires private providers, and if so, verifies the resolved config is // requires private providers, and if so, verifies the resolved config is
// marked as private. Returns nil if allowed, error if blocked. // marked as private. Returns nil if allowed, error if blocked.
func enforcePrivateProviderPolicy(userID, configID string) error { func enforcePrivateProviderPolicy(ctx context.Context, stores store.Stores, userID, configID string) error {
if configID == "" { if configID == "" {
return nil return nil
} }
// Check if user belongs to any team with require_private_providers policy requiresPrivate, err := stores.Teams.HasPrivateProviderRequirement(ctx, userID)
var requiresPrivate bool
var query string
if database.IsSQLite() {
query = `
SELECT EXISTS(
SELECT 1 FROM team_members tm
JOIN teams t ON t.id = tm.team_id
WHERE tm.user_id = ?
AND t.is_active = 1
AND json_extract(t.settings, '$.require_private_providers') = 'true'
)`
} else {
query = `
SELECT EXISTS(
SELECT 1 FROM team_members tm
JOIN teams t ON t.id = tm.team_id
WHERE tm.user_id = $1
AND t.is_active = true
AND t.settings->>'require_private_providers' = 'true'
)`
}
err := database.DB.QueryRow(query, userID).Scan(&requiresPrivate)
if err != nil || !requiresPrivate { if err != nil || !requiresPrivate {
return nil return nil
} }
// User is in a restricted team — verify the config is private // User is in a restricted team — verify the config is private
var isPrivate bool cfg, err := stores.Providers.GetByID(ctx, configID)
err = database.DB.QueryRow(database.Q(`
SELECT COALESCE(is_private, false) FROM provider_configs WHERE id = $1
`), configID).Scan(&isPrivate)
if err != nil { if err != nil {
return nil // config lookup failed, allow (fail open) return nil // config lookup failed, allow (fail open)
} }
if !isPrivate { if !cfg.IsPrivate {
return fmt.Errorf("your team requires private providers — this provider sends data externally") return fmt.Errorf("your team requires private providers — this provider sends data externally")
} }
return nil return nil
@@ -561,84 +471,31 @@ func (h *TeamHandler) ListTeamAuditLog(c *gin.Context) {
teamID := c.Param("teamId") teamID := c.Param("teamId")
page, perPage, offset := parsePagination(c) page, perPage, offset := parsePagination(c)
// Build filter clauses — always scoped to team members. opts := store.AuditListOptions{
// Use ? placeholders and convert for Postgres if needed. ListOptions: store.ListOptions{
clauses := []string{"al.actor_id IN (SELECT user_id FROM team_members WHERE team_id = ?)"} Limit: perPage,
args := []interface{}{teamID} Offset: offset,
},
TeamID: teamID,
}
if action := c.Query("action"); action != "" { if action := c.Query("action"); action != "" {
clauses = append(clauses, "al.action = ?") opts.Action = action
args = append(args, action)
} }
if actorID := c.Query("actor_id"); actorID != "" { if actorID := c.Query("actor_id"); actorID != "" {
clauses = append(clauses, "al.actor_id = ?") opts.ActorID = actorID
args = append(args, actorID)
} }
if rt := c.Query("resource_type"); rt != "" { if rt := c.Query("resource_type"); rt != "" {
clauses = append(clauses, "al.resource_type = ?") opts.ResourceType = rt
args = append(args, rt)
} }
where := "WHERE " + strings.Join(clauses, " AND ") entries, total, err := h.stores.Audit.List(c.Request.Context(), opts)
// For Postgres, convert ? to $N
if database.IsPostgres() {
where = convertPlaceholders(where)
}
// Count
var total int
countArgs := make([]interface{}, len(args))
copy(countArgs, args)
err := database.DB.QueryRow(`SELECT COUNT(*) FROM audit_log al `+where, countArgs...).Scan(&total)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "count failed"})
return
}
// Query with actor name join
limitOffset := fmt.Sprintf("LIMIT %d OFFSET %d", perPage, offset)
query := `
SELECT al.id, al.actor_id, COALESCE(u.username, '') as actor_name,
al.action, al.resource_type, al.resource_id,
COALESCE(al.metadata, '{}'), al.ip_address, al.created_at
FROM audit_log al
LEFT JOIN users u ON al.actor_id = u.id
` + where + `
ORDER BY al.created_at DESC
` + limitOffset
rows, err := database.DB.Query(query, args...)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "query failed"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "query failed"})
return return
} }
defer rows.Close() if entries == nil {
entries = []models.AuditEntry{}
type entry struct {
ID string `json:"id"`
ActorID *string `json:"actor_id"`
ActorName *string `json:"actor_name"`
Action string `json:"action"`
ResourceType string `json:"resource_type"`
ResourceID *string `json:"resource_id"`
Metadata string `json:"metadata"`
IPAddress *string `json:"ip_address"`
CreatedAt time.Time `json:"created_at"`
}
entries := make([]entry, 0)
for rows.Next() {
var e entry
var actorName sql.NullString
if err := rows.Scan(&e.ID, &e.ActorID, &actorName, &e.Action,
&e.ResourceType, &e.ResourceID, &e.Metadata, &e.IPAddress, database.ST(&e.CreatedAt)); err != nil {
continue
}
if actorName.Valid {
e.ActorName = &actorName.String
}
entries = append(entries, e)
} }
c.JSON(http.StatusOK, gin.H{ c.JSON(http.StatusOK, gin.H{
@@ -652,39 +509,11 @@ func (h *TeamHandler) ListTeamAuditLog(c *gin.Context) {
func (h *TeamHandler) ListTeamAuditActions(c *gin.Context) { func (h *TeamHandler) ListTeamAuditActions(c *gin.Context) {
teamID := c.Param("teamId") teamID := c.Param("teamId")
rows, err := database.DB.Query(database.Q(` actions, err := h.stores.Teams.ListTeamAuditActions(c.Request.Context(), teamID)
SELECT DISTINCT al.action
FROM audit_log al
WHERE al.actor_id IN (SELECT user_id FROM team_members WHERE team_id = $1)
ORDER BY al.action ASC
`), teamID)
if err != nil { if err != nil {
c.JSON(http.StatusOK, gin.H{"actions": []string{}}) c.JSON(http.StatusOK, gin.H{"actions": []string{}})
return return
} }
defer rows.Close()
actions := make([]string, 0)
for rows.Next() {
var a string
if rows.Scan(&a) == nil {
actions = append(actions, a)
}
}
c.JSON(http.StatusOK, gin.H{"actions": actions}) c.JSON(http.StatusOK, gin.H{"actions": actions})
} }
// convertPlaceholders converts ? placeholders to $1, $2, etc. for Postgres.
func convertPlaceholders(q string) string {
n := 1
var result strings.Builder
for _, ch := range q {
if ch == '?' {
result.WriteString(fmt.Sprintf("$%d", n))
n++
} else {
result.WriteRune(ch)
}
}
return result.String()
}

View File

@@ -1,5 +1,9 @@
package handlers package handlers
// workflow_assignments.go — Assignment queue for human review stages.
//
// v0.29.0: Raw SQL replaced with WorkflowStore + ChannelStore methods.
import ( import (
"encoding/json" "encoding/json"
"net/http" "net/http"
@@ -7,104 +11,51 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/events" "git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/notifications" "git.gobha.me/xcaliber/chat-switchboard/notifications"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
// ── Workflow Assignment Handler ───────────── // ── Workflow Assignment Handler ─────────────
// Manages the assignment queue for human review stages.
type WorkflowAssignmentHandler struct { type WorkflowAssignmentHandler struct {
hub *events.Hub stores store.Stores
hub *events.Hub
} }
func NewWorkflowAssignmentHandler(hub ...*events.Hub) *WorkflowAssignmentHandler { func NewWorkflowAssignmentHandler(stores store.Stores, hub ...*events.Hub) *WorkflowAssignmentHandler {
h := &WorkflowAssignmentHandler{} h := &WorkflowAssignmentHandler{stores: stores}
if len(hub) > 0 { if len(hub) > 0 {
h.hub = hub[0] h.hub = hub[0]
} }
return h return h
} }
type assignmentRow struct {
ID string `json:"id"`
ChannelID string `json:"channel_id"`
Stage int `json:"stage"`
TeamID string `json:"team_id"`
AssignedTo *string `json:"assigned_to"`
Status string `json:"status"`
CreatedAt time.Time `json:"created_at"`
ClaimedAt *time.Time `json:"claimed_at"`
CompletedAt *time.Time `json:"completed_at"`
}
// ListForTeam returns unassigned + claimed assignments for a team. // ListForTeam returns unassigned + claimed assignments for a team.
// GET /api/v1/teams/:teamId/assignments // GET /api/v1/teams/:teamId/assignments
func (h *WorkflowAssignmentHandler) ListForTeam(c *gin.Context) { func (h *WorkflowAssignmentHandler) ListForTeam(c *gin.Context) {
teamID := c.Param("teamId") teamID := c.Param("teamId")
status := c.DefaultQuery("status", "unassigned") status := c.DefaultQuery("status", "unassigned")
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(` result, err := h.stores.Workflows.ListAssignmentsForTeam(c.Request.Context(), teamID, status)
SELECT id, channel_id, stage, team_id, assigned_to, status,
created_at, claimed_at, completed_at
FROM workflow_assignments
WHERE team_id = $1 AND status = $2
ORDER BY created_at ASC
`), teamID, status)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list assignments"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list assignments"})
return return
} }
defer rows.Close()
var result []assignmentRow
for rows.Next() {
var a assignmentRow
if err := rows.Scan(&a.ID, &a.ChannelID, &a.Stage, &a.TeamID,
&a.AssignedTo, &a.Status, database.ST(&a.CreatedAt), database.SNT(&a.ClaimedAt), database.SNT(&a.CompletedAt)); err != nil {
continue
}
result = append(result, a)
}
if result == nil {
result = []assignmentRow{}
}
c.JSON(http.StatusOK, gin.H{"data": result}) c.JSON(http.StatusOK, gin.H{"data": result})
} }
// ListForUser returns assignments claimed by the current user plus // ListMine returns assignments claimed by the current user plus
// unassigned assignments for teams the user belongs to. // unassigned assignments for teams the user belongs to.
// GET /api/v1/workflow-assignments/mine // GET /api/v1/workflow-assignments/mine
func (h *WorkflowAssignmentHandler) ListMine(c *gin.Context) { func (h *WorkflowAssignmentHandler) ListMine(c *gin.Context) {
userID := c.GetString("user_id") userID := c.GetString("user_id")
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(`
SELECT DISTINCT wa.id, wa.channel_id, wa.stage, wa.team_id, wa.assigned_to, wa.status, result, err := h.stores.Workflows.ListAssignmentsMine(c.Request.Context(), userID)
wa.created_at, wa.claimed_at, wa.completed_at
FROM workflow_assignments wa
LEFT JOIN team_members tm ON tm.team_id = wa.team_id AND tm.user_id = $1
WHERE (wa.assigned_to = $2 AND wa.status = 'claimed')
OR (wa.status = 'unassigned' AND tm.user_id IS NOT NULL)
ORDER BY wa.created_at DESC
`), userID, userID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list assignments"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to list assignments"})
return return
} }
defer rows.Close()
var result []assignmentRow
for rows.Next() {
var a assignmentRow
if err := rows.Scan(&a.ID, &a.ChannelID, &a.Stage, &a.TeamID,
&a.AssignedTo, &a.Status, database.ST(&a.CreatedAt), database.SNT(&a.ClaimedAt), database.SNT(&a.CompletedAt)); err != nil {
continue
}
result = append(result, a)
}
if result == nil {
result = []assignmentRow{}
}
c.JSON(http.StatusOK, gin.H{"data": result}) c.JSON(http.StatusOK, gin.H{"data": result})
} }
@@ -115,30 +66,20 @@ func (h *WorkflowAssignmentHandler) Claim(c *gin.Context) {
userID := c.GetString("user_id") userID := c.GetString("user_id")
now := time.Now().UTC() now := time.Now().UTC()
res, err := database.DB.ExecContext(c.Request.Context(), database.Q(` n, err := h.stores.Workflows.ClaimAssignment(c.Request.Context(), assignmentID, userID)
UPDATE workflow_assignments
SET assigned_to = $1, status = 'claimed', claimed_at = $2
WHERE id = $3 AND status = 'unassigned'
`), userID, now, assignmentID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to claim assignment"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to claim assignment"})
return return
} }
n, _ := res.RowsAffected()
if n == 0 { if n == 0 {
c.JSON(http.StatusConflict, gin.H{"error": "assignment already claimed or not found"}) c.JSON(http.StatusConflict, gin.H{"error": "assignment already claimed or not found"})
return return
} }
// Look up channel for this assignment (needed for WS delivery + notification) // Look up channel for WS delivery + notification
var channelID string channelID, _ := h.stores.Workflows.GetAssignmentChannelID(c.Request.Context(), assignmentID)
_ = database.DB.QueryRowContext(c.Request.Context(),
database.Q(`SELECT channel_id FROM workflow_assignments WHERE id = $1`),
assignmentID).Scan(&channelID)
// Emit workflow.claimed WS event to all user participants in the channel. // Emit workflow.claimed WS event to all user participants in the channel
// Uses SendToUser (not room-scoped Bus.Publish) because room subscriptions
// are not yet wired on the client side. See websocket.md § Room Model.
if h.hub != nil && channelID != "" { if h.hub != nil && channelID != "" {
payload, _ := json.Marshal(map[string]any{ payload, _ := json.Marshal(map[string]any{
"assignment_id": assignmentID, "assignment_id": assignmentID,
@@ -150,22 +91,13 @@ func (h *WorkflowAssignmentHandler) Claim(c *gin.Context) {
Payload: payload, Payload: payload,
Ts: now.UnixMilli(), Ts: now.UnixMilli(),
} }
rows, err := database.DB.QueryContext(c.Request.Context(), database.Q(` pids, _ := h.stores.Channels.ListUserParticipantIDs(c.Request.Context(), channelID, "")
SELECT participant_id FROM channel_participants for _, uid := range pids {
WHERE channel_id = $1 AND participant_type = 'user' h.hub.SendToUser(uid, evt)
`), channelID)
if err == nil {
defer rows.Close()
for rows.Next() {
var uid string
if rows.Scan(&uid) == nil {
h.hub.SendToUser(uid, evt)
}
}
} }
} }
// Persist notification for bell/inbox (v0.28.2) // Persist notification for bell/inbox
if svc := notifications.Default(); svc != nil && channelID != "" { if svc := notifications.Default(); svc != nil && channelID != "" {
notifications.NotifyWorkflowClaimed(svc, userID, assignmentID, channelID) notifications.NotifyWorkflowClaimed(svc, userID, assignmentID, channelID)
} }
@@ -177,18 +109,12 @@ func (h *WorkflowAssignmentHandler) Claim(c *gin.Context) {
// POST /api/v1/workflow-assignments/:id/complete // POST /api/v1/workflow-assignments/:id/complete
func (h *WorkflowAssignmentHandler) Complete(c *gin.Context) { func (h *WorkflowAssignmentHandler) Complete(c *gin.Context) {
assignmentID := c.Param("id") assignmentID := c.Param("id")
now := time.Now().UTC()
res, err := database.DB.ExecContext(c.Request.Context(), database.Q(` n, err := h.stores.Workflows.CompleteAssignment(c.Request.Context(), assignmentID)
UPDATE workflow_assignments
SET status = 'completed', completed_at = $1
WHERE id = $2 AND status = 'claimed'
`), now, assignmentID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to complete assignment"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to complete assignment"})
return return
} }
n, _ := res.RowsAffected()
if n == 0 { if n == 0 {
c.JSON(http.StatusConflict, gin.H{"error": "assignment not in claimed state"}) c.JSON(http.StatusConflict, gin.H{"error": "assignment not in claimed state"})
return return

View File

@@ -4,11 +4,9 @@ import (
"fmt" "fmt"
"log" "log"
"net/http" "net/http"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
@@ -77,20 +75,15 @@ func (h *WorkflowEntryHandler) StartVisitor(c *gin.Context) {
} }
// Set workflow columns // Set workflow columns
allowAnonVal := interface{}(true) err = h.stores.Channels.SetWorkflowInstance(ctx, ch.ID, wf.ID, ver.VersionNumber, []byte("{}"), "active")
if database.CurrentDialect == database.DialectSQLite {
allowAnonVal = 1
}
_, err = database.DB.ExecContext(ctx, database.Q(`
UPDATE channels
SET workflow_id = $1, workflow_version = $2, current_stage = 0,
stage_data = '{}', workflow_status = 'active',
last_activity_at = $3, allow_anonymous = $4, ai_mode = 'auto'
WHERE id = $5
`), wf.ID, ver.VersionNumber, time.Now().UTC(), allowAnonVal, ch.ID)
if err != nil { if err != nil {
log.Printf("Failed to set workflow columns: %v", err) log.Printf("Failed to set workflow columns: %v", err)
} }
// Enable anonymous access + auto AI mode for visitor entry
_ = h.stores.Channels.Update(ctx, ch.ID, map[string]interface{}{
"allow_anonymous": true,
"ai_mode": "auto",
})
// Create anonymous session // Create anonymous session
sessionToken := store.NewID() sessionToken := store.NewID()

View File

@@ -79,7 +79,7 @@ func setupWorkflowInstanceHarness(t *testing.T) *workflowInstanceHarness {
protected.POST("/channels/:id/workflow/reject", wfInstH.Reject) protected.POST("/channels/:id/workflow/reject", wfInstH.Reject)
// Assignments // Assignments
wfAssignH := NewWorkflowAssignmentHandler() wfAssignH := NewWorkflowAssignmentHandler(stores)
protected.GET("/workflow-assignments/mine", wfAssignH.ListMine) protected.GET("/workflow-assignments/mine", wfAssignH.ListMine)
protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim) protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim)
protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete) protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete)
@@ -93,7 +93,7 @@ func setupWorkflowInstanceHarness(t *testing.T) *workflowInstanceHarness {
// Team-scoped assignment listing (mirrors main.go teamScoped) // Team-scoped assignment listing (mirrors main.go teamScoped)
teamScoped := protected.Group("/teams/:teamId") teamScoped := protected.Group("/teams/:teamId")
teamAssignH := NewWorkflowAssignmentHandler() teamAssignH := NewWorkflowAssignmentHandler(stores)
teamScoped.GET("/assignments", teamAssignH.ListForTeam) teamScoped.GET("/assignments", teamAssignH.ListForTeam)
// Visitor entry // Visitor entry

View File

@@ -3,13 +3,13 @@ package handlers
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"log" "log"
"net/http" "net/http"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/events" "git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/notifications" "git.gobha.me/xcaliber/chat-switchboard/notifications"
@@ -77,24 +77,14 @@ func (h *WorkflowInstanceHandler) Start(c *gin.Context) {
// Set workflow-specific columns (not part of base Channel.Create) // Set workflow-specific columns (not part of base Channel.Create)
allowAnon := wf.EntryMode == "public_link" allowAnon := wf.EntryMode == "public_link"
var allowAnonVal interface{} = allowAnon err = h.stores.Channels.SetWorkflowInstance(ctx, ch.ID, wfID, ver.VersionNumber, []byte("{}"), "active")
if database.CurrentDialect == database.DialectSQLite {
if allowAnon {
allowAnonVal = 1
} else {
allowAnonVal = 0
}
}
_, err = database.DB.ExecContext(ctx, database.Q(`
UPDATE channels
SET workflow_id = $1, workflow_version = $2, current_stage = 0,
stage_data = '{}', workflow_status = 'active',
last_activity_at = $3, allow_anonymous = $4, ai_mode = 'auto'
WHERE id = $5
`), wfID, ver.VersionNumber, time.Now().UTC(), allowAnonVal, ch.ID)
if err != nil { if err != nil {
log.Printf("Failed to set workflow columns on channel %s: %v", ch.ID, err) log.Printf("Failed to set workflow columns on channel %s: %v", ch.ID, err)
} }
_ = h.stores.Channels.Update(ctx, ch.ID, map[string]interface{}{
"allow_anonymous": allowAnon,
"ai_mode": "auto",
})
// Add caller as channel owner // Add caller as channel owner
_ = h.stores.Channels.AddParticipant(ctx, &models.ChannelParticipant{ _ = h.stores.Channels.AddParticipant(ctx, &models.ChannelParticipant{
@@ -126,34 +116,15 @@ func (h *WorkflowInstanceHandler) Start(c *gin.Context) {
// ── Status ────────────────────────────────── // ── Status ──────────────────────────────────
// WorkflowChannelStatus holds the runtime state of a workflow instance.
type WorkflowChannelStatus struct {
WorkflowID *string `json:"workflow_id"`
WorkflowVersion *int `json:"workflow_version"`
CurrentStage int `json:"current_stage"`
StageData json.RawMessage `json:"stage_data"`
Status string `json:"status"`
LastActivityAt *string `json:"last_activity_at"`
}
// GetStatus returns the workflow state for a channel. // GetStatus returns the workflow state for a channel.
// GET /api/v1/channels/:id/workflow/status // GET /api/v1/channels/:id/workflow/status
func (h *WorkflowInstanceHandler) GetStatus(c *gin.Context) { func (h *WorkflowInstanceHandler) GetStatus(c *gin.Context) {
channelID := c.Param("id") channelID := c.Param("id")
var ws WorkflowChannelStatus ws, err := h.stores.Channels.GetWorkflowStatus(c.Request.Context(), channelID)
var stageData []byte if err != nil || ws == nil {
err := database.DB.QueryRowContext(c.Request.Context(), database.Q(`
SELECT workflow_id, workflow_version, current_stage,
COALESCE(stage_data, '{}'), COALESCE(workflow_status, 'active'),
last_activity_at
FROM channels WHERE id = $1 AND type = 'workflow'
`), channelID).Scan(&ws.WorkflowID, &ws.WorkflowVersion,
&ws.CurrentStage, &stageData, &ws.Status, &ws.LastActivityAt)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "workflow channel not found"}) c.JSON(http.StatusNotFound, gin.H{"error": "workflow channel not found"})
return return
} }
ws.StageData = stageData
c.JSON(http.StatusOK, ws) c.JSON(http.StatusOK, ws)
} }
@@ -186,17 +157,12 @@ func (h *WorkflowInstanceHandler) Advance(c *gin.Context) {
} }
_ = c.ShouldBindJSON(&body) _ = c.ShouldBindJSON(&body)
mergedData := tools.MergeWorkflowStageData(ctx, channelID, body.Data) mergedData := tools.MergeWorkflowStageData(ctx, h.stores.Channels, channelID, body.Data)
nextStage := currentStage + 1 nextStage := currentStage + 1
if nextStage >= len(stages) { if nextStage >= len(stages) {
// Workflow complete // Workflow complete
_, err = database.DB.ExecContext(ctx, database.Q(` err = h.stores.Channels.CompleteWorkflow(ctx, channelID, nextStage, json.RawMessage(mergedData))
UPDATE channels
SET current_stage = $1, workflow_status = 'completed',
stage_data = $2, last_activity_at = $3, ai_mode = 'off'
WHERE id = $4
`), nextStage, mergedData, time.Now().UTC(), channelID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to complete workflow"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to complete workflow"})
return return
@@ -216,11 +182,7 @@ func (h *WorkflowInstanceHandler) Advance(c *gin.Context) {
} }
// Advance to next stage // Advance to next stage
_, err = database.DB.ExecContext(ctx, database.Q(` err = h.stores.Channels.AdvanceWorkflowStage(ctx, channelID, nextStage, json.RawMessage(mergedData))
UPDATE channels
SET current_stage = $1, stage_data = $2, last_activity_at = $3
WHERE id = $4
`), nextStage, mergedData, time.Now().UTC(), channelID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to advance stage"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to advance stage"})
return return
@@ -245,7 +207,7 @@ func (h *WorkflowInstanceHandler) Advance(c *gin.Context) {
// v0.27.0: Assignment + round-robin + WS notifications // v0.27.0: Assignment + round-robin + WS notifications
if nextStageDef.AssignmentTeamID != nil { if nextStageDef.AssignmentTeamID != nil {
assignmentID := tools.CreateWorkflowAssignment(ctx, channelID, nextStage, *nextStageDef.AssignmentTeamID) assignmentID := tools.CreateWorkflowAssignment(ctx, h.stores, channelID, nextStage, *nextStageDef.AssignmentTeamID)
// Round-robin auto-assignment if configured // Round-robin auto-assignment if configured
assignedTo := h.tryRoundRobin(ctx, nextStageDef, assignmentID) assignedTo := h.tryRoundRobin(ctx, nextStageDef, assignmentID)
@@ -298,9 +260,7 @@ func (h *WorkflowInstanceHandler) Reject(c *gin.Context) {
} }
prevStage := currentStage - 1 prevStage := currentStage - 1
_, err = database.DB.ExecContext(ctx, database.Q(` err = h.stores.Channels.RejectWorkflowToStage(ctx, channelID, prevStage)
UPDATE channels SET current_stage = $1, last_activity_at = $2 WHERE id = $3
`), prevStage, time.Now().UTC(), channelID)
if err != nil { if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to reject stage"}) c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to reject stage"})
return return
@@ -327,15 +287,14 @@ func (h *WorkflowInstanceHandler) Reject(c *gin.Context) {
// ── Helpers ───────────────────────────────── // ── Helpers ─────────────────────────────────
func (h *WorkflowInstanceHandler) readWorkflowState(ctx context.Context, channelID string) (workflowID string, currentStage int, status string, err error) { func (h *WorkflowInstanceHandler) readWorkflowState(ctx context.Context, channelID string) (workflowID string, currentStage int, status string, err error) {
var wfID *string ws, err := h.stores.Channels.GetWorkflowStatus(ctx, channelID)
err = database.DB.QueryRowContext(ctx, database.Q(` if err != nil || ws == nil {
SELECT workflow_id, COALESCE(current_stage, 0), COALESCE(workflow_status, 'active') return "", 0, "", fmt.Errorf("workflow channel not found")
FROM channels WHERE id = $1 AND type = 'workflow'
`), channelID).Scan(&wfID, &currentStage, &status)
if wfID != nil {
workflowID = *wfID
} }
return if ws.WorkflowID != nil {
workflowID = *ws.WorkflowID
}
return workflowID, ws.CurrentStage, ws.Status, nil
} }
// emitWorkflowEvent pushes a workflow event to all user participants in the channel. // emitWorkflowEvent pushes a workflow event to all user participants in the channel.
@@ -353,20 +312,13 @@ func (h *WorkflowInstanceHandler) emitWorkflowEvent(label, channelID string, dat
} }
// Send to all user participants in the channel // Send to all user participants in the channel
rows, err := database.DB.Query(database.Q(` pids, err := h.stores.Channels.ListUserParticipantIDs(context.Background(), channelID, "")
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'user'
`), channelID)
if err != nil { if err != nil {
log.Printf("[ws] %s: failed to query participants for channel %s: %v", label, channelID[:min(8, len(channelID))], err) log.Printf("[ws] %s: failed to query participants for channel %s: %v", label, channelID[:min(8, len(channelID))], err)
return return
} }
defer rows.Close() for _, uid := range pids {
for rows.Next() { h.hub.SendToUser(uid, evt)
var uid string
if rows.Scan(&uid) == nil {
h.hub.SendToUser(uid, evt)
}
} }
} }
@@ -396,52 +348,15 @@ func (h *WorkflowInstanceHandler) tryRoundRobin(ctx context.Context, stage model
return "" return ""
} }
// Get team members assignedTo, err := h.stores.Workflows.TryRoundRobin(ctx, *stage.AssignmentTeamID, assignmentID)
members, err := h.stores.Teams.ListMembers(ctx, *stage.AssignmentTeamID)
if err != nil || len(members) == 0 {
return ""
}
// Find the least-recently-assigned member.
// Query: for each member, find their most recent claimed_at in workflow_assignments.
// Pick the member with the oldest (or null) claimed_at.
var bestUserID string
bestUserID = members[0].UserID // fallback to first member
rows, err := database.DB.QueryContext(ctx, database.Q(`
SELECT m.user_id, COALESCE(MAX(wa.claimed_at), '1970-01-01T00:00:00Z') as last_claim
FROM team_members m
LEFT JOIN workflow_assignments wa ON wa.assigned_to = m.user_id AND wa.team_id = $1
WHERE m.team_id = $2
GROUP BY m.user_id
ORDER BY last_claim ASC
LIMIT 1
`), *stage.AssignmentTeamID, *stage.AssignmentTeamID)
if err == nil {
defer rows.Close()
if rows.Next() {
var uid string
var lastClaim string // COALESCE returns TEXT on both dialects
if err := rows.Scan(&uid, &lastClaim); err == nil {
bestUserID = uid
}
}
}
// Claim the assignment for this user
now := time.Now().UTC()
_, err = database.DB.ExecContext(ctx, database.Q(`
UPDATE workflow_assignments
SET assigned_to = $1, status = 'claimed', claimed_at = $2
WHERE id = $3 AND status = 'unassigned'
`), bestUserID, now, assignmentID)
if err != nil { if err != nil {
log.Printf("[workflow] round-robin: failed to auto-assign %s to %s: %v", assignmentID, bestUserID, err) log.Printf("[workflow] round-robin: failed to auto-assign %s: %v", assignmentID, err)
return "" return ""
} }
if assignedTo != "" {
log.Printf("[workflow] round-robin: auto-assigned %s to user %s", assignmentID, bestUserID) log.Printf("[workflow] round-robin: auto-assigned %s to user %s", assignmentID, assignedTo)
return bestUserID }
return assignedTo
} }
// notifyAssignment sends notifications to team members about a new workflow assignment. // notifyAssignment sends notifications to team members about a new workflow assignment.

View File

@@ -319,7 +319,7 @@ func setupWorkflowHarness(t *testing.T) *workflowHarness {
protected.POST("/channels/:id/workflow/reject", wfInstH.Reject) protected.POST("/channels/:id/workflow/reject", wfInstH.Reject)
// Channels (needed for workflow instance creation) // Channels (needed for workflow instance creation)
channels := NewChannelHandler() channels := NewChannelHandler(stores)
protected.GET("/channels", channels.ListChannels) protected.GET("/channels", channels.ListChannels)
protected.POST("/channels", channels.CreateChannel) protected.POST("/channels", channels.CreateChannel)
protected.GET("/channels/:id", channels.GetChannel) protected.GET("/channels/:id", channels.GetChannel)

View File

@@ -20,6 +20,8 @@ import (
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/events" "git.gobha.me/xcaliber/chat-switchboard/events"
"git.gobha.me/xcaliber/chat-switchboard/extraction" "git.gobha.me/xcaliber/chat-switchboard/extraction"
"git.gobha.me/xcaliber/chat-switchboard/filters"
"git.gobha.me/xcaliber/chat-switchboard/sandbox"
"git.gobha.me/xcaliber/chat-switchboard/handlers" "git.gobha.me/xcaliber/chat-switchboard/handlers"
"git.gobha.me/xcaliber/chat-switchboard/health" "git.gobha.me/xcaliber/chat-switchboard/health"
"git.gobha.me/xcaliber/chat-switchboard/knowledge" "git.gobha.me/xcaliber/chat-switchboard/knowledge"
@@ -37,6 +39,7 @@ import (
sqliteStore "git.gobha.me/xcaliber/chat-switchboard/store/sqlite" sqliteStore "git.gobha.me/xcaliber/chat-switchboard/store/sqlite"
"git.gobha.me/xcaliber/chat-switchboard/tools" "git.gobha.me/xcaliber/chat-switchboard/tools"
"git.gobha.me/xcaliber/chat-switchboard/tools/search" "git.gobha.me/xcaliber/chat-switchboard/tools/search"
"git.gobha.me/xcaliber/chat-switchboard/treepath"
"git.gobha.me/xcaliber/chat-switchboard/workspace" "git.gobha.me/xcaliber/chat-switchboard/workspace"
) )
@@ -103,6 +106,10 @@ func main() {
stores = postgres.NewStores(database.DB) stores = postgres.NewStores(database.DB)
} }
// v0.29.0: Wire store layer into treepath for backward compat.
// New code should call stores.Messages.* directly.
treepath.Stores = &stores
// Provider health accumulator (v0.22.0) // Provider health accumulator (v0.22.0)
if database.IsSQLite() { if database.IsSQLite() {
healthStore = sqliteStore.NewHealthStore() healthStore = sqliteStore.NewHealthStore()
@@ -171,43 +178,24 @@ func main() {
// Staleness: mark idle active instances as stale // Staleness: mark idle active instances as stale
cutoff := time.Now().UTC().Add(-time.Duration(cfg.WorkflowStaleHours) * time.Hour) cutoff := time.Now().UTC().Add(-time.Duration(cfg.WorkflowStaleHours) * time.Hour)
res, err := database.DB.ExecContext(ctx, database.Q(` n, err := stores.Channels.MarkStaleWorkflows(ctx, cutoff)
UPDATE channels
SET workflow_status = 'stale'
WHERE type = 'workflow'
AND workflow_status = 'active'
AND last_activity_at < $1
`), cutoff)
if err != nil { if err != nil {
log.Printf("⚠ workflow staleness sweep failed: %v", err) log.Printf("⚠ workflow staleness sweep failed: %v", err)
} else if n, _ := res.RowsAffected(); n > 0 { } else if n > 0 {
log.Printf("🧹 workflows: marked %d instances as stale", n) log.Printf("🧹 workflows: marked %d instances as stale", n)
} }
// v0.27.0: Retention enforcement — delete completed workflow channels // v0.27.0: Retention enforcement — delete completed workflow channels
// where the parent workflow has retention.mode="delete" and // where the parent workflow has retention.mode="delete" and
// retention.delete_after_days has elapsed since completion. // retention.delete_after_days has elapsed since completion.
res2, err := database.DB.ExecContext(ctx, database.Q(` n2, err := stores.Channels.EnforceWorkflowRetention(ctx)
DELETE FROM channels
WHERE type = 'workflow'
AND workflow_status IN ('completed', 'archived')
AND workflow_id IS NOT NULL
AND last_activity_at < $1
AND workflow_id IN (
SELECT id FROM workflows
WHERE retention IS NOT NULL
AND retention->>'mode' = 'delete'
AND (retention->>'delete_after_days')::int > 0
AND channels.last_activity_at < now() - ((retention->>'delete_after_days')::int || ' days')::interval
)
`), cutoff)
if err != nil { if err != nil {
// SQLite doesn't support JSON operators — skip retention on SQLite // SQLite doesn't support JSON operators — skip retention on SQLite
if !database.IsSQLite() { if !database.IsSQLite() {
log.Printf("⚠ workflow retention enforcement failed: %v", err) log.Printf("⚠ workflow retention enforcement failed: %v", err)
} }
} else if n, _ := res2.RowsAffected(); n > 0 { } else if n2 > 0 {
log.Printf("🧹 workflows: deleted %d expired instances (retention policy)", n) log.Printf("🧹 workflows: deleted %d expired instances (retention policy)", n2)
} }
cancel() cancel()
@@ -367,6 +355,26 @@ func main() {
memScanner.Start() memScanner.Start()
defer memScanner.Stop() defer memScanner.Stop()
// ── Pre-completion filter chain (v0.29.0) ──
// Built-in filters register here. Starlark extension filters will
// register at package install time (CS3).
filterChain := filters.NewChain()
filterChain.Register(filters.NewKBInjectFilter(stores))
// ── Starlark Runner (v0.29.0 CS3) ──
// Sandboxed interpreter for extension scripts. Runner assembles
// modules based on granted permissions. Notifier attached below
// after notification service init.
starlarkRunner := sandbox.NewRunner(
sandbox.New(sandbox.DefaultConfig()),
stores,
)
// Discover and register active Starlark pre-completion filters
filters.DiscoverStarlarkFilters(context.Background(), filterChain, stores, starlarkRunner)
log.Printf(" 🔗 Pre-completion filter chain: %d filters", filterChain.Len())
r := gin.Default() r := gin.Default()
userCache := middleware.NewUserStatusCache() userCache := middleware.NewUserStatusCache()
r.Use(middleware.CORS(cfg)) r.Use(middleware.CORS(cfg))
@@ -404,6 +412,7 @@ func main() {
} }
notifSvc.StartCleanup() notifSvc.StartCleanup()
notifications.SetDefault(notifSvc) notifications.SetDefault(notifSvc)
starlarkRunner.SetNotifier(notifSvc)
// Subscribe to role.fallback events → generate notifications for admins // Subscribe to role.fallback events → generate notifications for admins
bus.Subscribe("role.fallback", notifications.RoleFallbackHandler(notifSvc, stores)) bus.Subscribe("role.fallback", notifications.RoleFallbackHandler(notifSvc, stores))
@@ -416,6 +425,7 @@ func main() {
scheduler.RegisterBuiltins() scheduler.RegisterBuiltins()
exec := scheduler.NewExecutor(stores, keyResolver, hub, healthAccum) exec := scheduler.NewExecutor(stores, keyResolver, hub, healthAccum)
exec.SetRunner(starlarkRunner)
taskSched := scheduler.New(stores, exec) taskSched := scheduler.New(stores, exec)
go taskSched.Run() go taskSched.Run()
log.Println(" ⏰ Task scheduler started (with executor)") log.Println(" ⏰ Task scheduler started (with executor)")
@@ -533,7 +543,7 @@ func main() {
}) })
// Channels // Channels
channels := handlers.NewChannelHandler() channels := handlers.NewChannelHandler(stores)
protected.GET("/channels", channels.ListChannels) protected.GET("/channels", channels.ListChannels)
protected.POST("/channels", middleware.RequirePermission(auth.PermChannelCreate, stores), channels.CreateChannel) protected.POST("/channels", middleware.RequirePermission(auth.PermChannelCreate, stores), channels.CreateChannel)
protected.GET("/channels/:id", channels.GetChannel) protected.GET("/channels/:id", channels.GetChannel)
@@ -547,19 +557,19 @@ func main() {
channelID := c.Param("id") channelID := c.Param("id")
// Resolve display name // Resolve display name
var displayName string var displayName string
_ = database.DB.QueryRow(database.Q(` user, err := stores.Users.GetByID(c.Request.Context(), userID)
SELECT COALESCE(display_name, username) FROM users WHERE id = $1 if err == nil && user != nil {
`), userID).Scan(&displayName) displayName = user.DisplayName
if displayName == "" {
displayName = user.Username
}
}
if displayName == "" { if displayName == "" {
displayName = userID[:8] displayName = userID[:8]
} }
// Broadcast to other user participants // Broadcast to other user participants
pRows, err := database.DB.Query(database.Q(` pids, err := stores.Channels.ListUserParticipantIDs(c.Request.Context(), channelID, userID)
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'user' AND participant_id != $2
`), channelID, userID)
if err == nil { if err == nil {
defer pRows.Close()
payload, _ := json.Marshal(map[string]any{ payload, _ := json.Marshal(map[string]any{
"channel_id": channelID, "channel_id": channelID,
"user_id": userID, "user_id": userID,
@@ -570,32 +580,30 @@ func main() {
Payload: payload, Payload: payload,
Ts: time.Now().UnixMilli(), Ts: time.Now().UnixMilli(),
} }
for pRows.Next() { for _, pid := range pids {
var pid string hub.SendToUser(pid, evt)
if pRows.Scan(&pid) == nil {
hub.SendToUser(pid, evt)
}
} }
} }
c.JSON(200, gin.H{"ok": true}) c.JSON(200, gin.H{"ok": true})
}) })
// Chat Folders (v0.23.1) // Chat Folders (v0.23.1)
folders := handlers.NewFolderHandler() folders := handlers.NewFolderHandler(stores)
protected.GET("/folders", folders.List) protected.GET("/folders", folders.List)
protected.POST("/folders", folders.Create) protected.POST("/folders", folders.Create)
protected.PUT("/folders/:id", folders.Update) protected.PUT("/folders/:id", folders.Update)
protected.DELETE("/folders/:id", folders.Delete) protected.DELETE("/folders/:id", folders.Delete)
// Presence (v0.23.1) // Presence (v0.23.1)
protected.POST("/presence/heartbeat", handlers.PresenceHeartbeat) presence := handlers.NewPresenceHandler(stores)
protected.GET("/presence", handlers.PresenceQuery) protected.POST("/presence/heartbeat", presence.Heartbeat)
protected.GET("/presence", presence.Query)
// User search (v0.23.2 — DM user picker) // User search (v0.23.2 — DM user picker)
protected.GET("/users/search", handlers.SearchUsers) protected.GET("/users/search", presence.SearchUsers)
// Persona groups (v0.23.2 — roster templates for group chats) // Persona groups (v0.23.2 — roster templates for group chats)
pgH := handlers.NewPersonaGroupHandler() pgH := handlers.NewPersonaGroupHandler(stores)
protected.GET("/persona-groups", pgH.List) protected.GET("/persona-groups", pgH.List)
protected.POST("/persona-groups", pgH.Create) protected.POST("/persona-groups", pgH.Create)
protected.GET("/persona-groups/:id", pgH.Get) protected.GET("/persona-groups/:id", pgH.Get)
@@ -627,7 +635,7 @@ func main() {
protected.POST("/channels/:id/workflow/reject", wfInstH.Reject) protected.POST("/channels/:id/workflow/reject", wfInstH.Reject)
// Workflow assignments (v0.26.4 — team assignment queue) // Workflow assignments (v0.26.4 — team assignment queue)
wfAssignH := handlers.NewWorkflowAssignmentHandler(hub) wfAssignH := handlers.NewWorkflowAssignmentHandler(stores, hub)
protected.GET("/workflow-assignments/mine", wfAssignH.ListMine) protected.GET("/workflow-assignments/mine", wfAssignH.ListMine)
protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim) protected.POST("/workflow-assignments/:id/claim", wfAssignH.Claim)
protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete) protected.POST("/workflow-assignments/:id/complete", wfAssignH.Complete)
@@ -676,6 +684,7 @@ func main() {
comp.SetHealthStore(healthStore) comp.SetHealthStore(healthStore)
} }
comp.SetRoutingEvaluator(routing.NewEvaluator()) comp.SetRoutingEvaluator(routing.NewEvaluator())
comp.SetFilterChain(filterChain)
protected.POST("/chat/completions", comp.Complete) protected.POST("/chat/completions", comp.Complete)
protected.GET("/tools", comp.ListTools) protected.GET("/tools", comp.ListTools)
@@ -685,7 +694,7 @@ func main() {
// Summarize & Continue (backed by compaction service) // Summarize & Continue (backed by compaction service)
compactionSvc := compaction.NewService(stores, roleResolver) compactionSvc := compaction.NewService(stores, roleResolver)
summarize := handlers.NewSummarizeHandler(compactionSvc) summarize := handlers.NewSummarizeHandler(stores, compactionSvc)
protected.POST("/channels/:id/summarize", summarize.Summarize) protected.POST("/channels/:id/summarize", summarize.Summarize)
// Auto-title generation (utility role) // Auto-title generation (utility role)
@@ -716,7 +725,7 @@ func main() {
protected.POST("/models/preferences/bulk", modelPrefs.BulkSetPreferences) protected.POST("/models/preferences/bulk", modelPrefs.BulkSetPreferences)
// User Settings & Profile // User Settings & Profile
settings := handlers.NewSettingsHandler(uekCache) settings := handlers.NewSettingsHandler(stores, uekCache)
protected.GET("/profile", settings.GetProfile) protected.GET("/profile", settings.GetProfile)
protected.PUT("/profile", settings.UpdateProfile) protected.PUT("/profile", settings.UpdateProfile)
protected.POST("/profile/password", settings.ChangePassword) protected.POST("/profile/password", settings.ChangePassword)
@@ -892,7 +901,7 @@ func main() {
// Team admin self-service // Team admin self-service
teamScoped := protected.Group("/teams/:teamId") teamScoped := protected.Group("/teams/:teamId")
teamScoped.Use(middleware.RequireTeamAdmin()) teamScoped.Use(middleware.RequireTeamAdmin(stores.Teams))
{ {
teamScoped.GET("/members", teams.ListMembers) teamScoped.GET("/members", teams.ListMembers)
teamScoped.POST("/members", teams.AddMember) teamScoped.POST("/members", teams.AddMember)
@@ -938,7 +947,7 @@ func main() {
teamScoped.DELETE("/roles/:role", teamRoles.DeleteTeamRole) teamScoped.DELETE("/roles/:role", teamRoles.DeleteTeamRole)
// Team workflow assignments (v0.26.4) // Team workflow assignments (v0.26.4)
teamAssignH := handlers.NewWorkflowAssignmentHandler(hub) teamAssignH := handlers.NewWorkflowAssignmentHandler(stores, hub)
teamScoped.GET("/assignments", teamAssignH.ListForTeam) teamScoped.GET("/assignments", teamAssignH.ListForTeam)
// Team tasks — admin CRUD (v0.27.5) // Team tasks — admin CRUD (v0.27.5)
@@ -952,7 +961,7 @@ func main() {
// Team task viewing for all members (v0.27.5) // Team task viewing for all members (v0.27.5)
teamMemberRoutes := protected.Group("/teams/:teamId") teamMemberRoutes := protected.Group("/teams/:teamId")
teamMemberRoutes.Use(middleware.RequireTeamMember()) teamMemberRoutes.Use(middleware.RequireTeamMember(stores.Teams))
{ {
teamMemberTaskH := handlers.NewTaskHandler(stores) teamMemberTaskH := handlers.NewTaskHandler(stores)
teamMemberRoutes.GET("/tasks", teamMemberTaskH.ListTeamTasks) teamMemberRoutes.GET("/tasks", teamMemberTaskH.ListTeamTasks)
@@ -1015,8 +1024,8 @@ func main() {
admin.POST("/personas", personaAdm.CreateAdminPersona) admin.POST("/personas", personaAdm.CreateAdminPersona)
admin.PUT("/personas/:id", personaAdm.UpdateAdminPersona) admin.PUT("/personas/:id", personaAdm.UpdateAdminPersona)
admin.DELETE("/personas/:id", personaAdm.DeleteAdminPersona) admin.DELETE("/personas/:id", personaAdm.DeleteAdminPersona)
admin.POST("/personas/:id/avatar", handlers.UploadPersonaAvatar) admin.POST("/personas/:id/avatar", func(c *gin.Context) { handlers.UploadPersonaAvatar(stores.Personas, c) })
admin.DELETE("/personas/:id/avatar", handlers.DeletePersonaAvatar) admin.DELETE("/personas/:id/avatar", func(c *gin.Context) { handlers.DeletePersonaAvatar(stores.Personas, c) })
admin.GET("/personas/:id/knowledge-bases", personaAdm.GetPersonaKBs) // v0.17.0 admin.GET("/personas/:id/knowledge-bases", personaAdm.GetPersonaKBs) // v0.17.0
admin.PUT("/personas/:id/knowledge-bases", personaAdm.SetPersonaKBs) // v0.17.0 admin.PUT("/personas/:id/knowledge-bases", personaAdm.SetPersonaKBs) // v0.17.0
admin.GET("/personas/:id/tool-grants", personaAdm.GetPersonaToolGrants) // v0.25.0 admin.GET("/personas/:id/tool-grants", personaAdm.GetPersonaToolGrants) // v0.25.0
@@ -1116,6 +1125,20 @@ func main() {
admin.PUT("/extensions/:id", extAdm.AdminUpdateExtension) admin.PUT("/extensions/:id", extAdm.AdminUpdateExtension)
admin.DELETE("/extensions/:id", extAdm.AdminUninstallExtension) admin.DELETE("/extensions/:id", extAdm.AdminUninstallExtension)
// Extension permissions (admin — v0.29.0)
extPermH := handlers.NewExtPermHandler(stores)
admin.GET("/extensions/:id/permissions", extPermH.ListPackagePermissions)
admin.GET("/extensions/:id/review", extPermH.ReviewPackage)
admin.POST("/extensions/:id/permissions/:perm/grant", extPermH.GrantPermission)
admin.POST("/extensions/:id/permissions/:perm/revoke", extPermH.RevokePermission)
admin.POST("/extensions/:id/permissions/grant-all", extPermH.GrantAllPermissions)
// Extension secrets (admin — v0.29.0 CS3)
extSecH := handlers.NewExtSecretsHandler(stores)
admin.GET("/extensions/:id/secrets", extSecH.GetSecrets)
admin.PUT("/extensions/:id/secrets", extSecH.SetSecrets)
admin.DELETE("/extensions/:id/secrets", extSecH.DeleteSecrets)
// Provider Health (admin — v0.22.0) // Provider Health (admin — v0.22.0)
healthAdm := handlers.NewHealthAdminHandler(healthStore, stores) healthAdm := handlers.NewHealthAdminHandler(healthStore, stores)
admin.GET("/providers/health", healthAdm.GetAllProviderHealth) admin.GET("/providers/health", healthAdm.GetAllProviderHealth)
@@ -1222,6 +1245,7 @@ func main() {
wfAPI.GET("/:id/messages", wfMsgs.ListMessages) wfAPI.GET("/:id/messages", wfMsgs.ListMessages)
wfComp := handlers.NewCompletionHandler(keyResolver, stores, hub, objStore, kbEmbedder) wfComp := handlers.NewCompletionHandler(keyResolver, stores, hub, objStore, kbEmbedder)
wfComp.SetFilterChain(filterChain)
wfAPI.POST("/:id/completions", wfComp.Complete) wfAPI.POST("/:id/completions", wfComp.Complete)
} }

View File

@@ -7,7 +7,6 @@ import (
"log" "log"
"strings" "strings"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/knowledge" "git.gobha.me/xcaliber/chat-switchboard/knowledge"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/notifications" "git.gobha.me/xcaliber/chat-switchboard/notifications"
@@ -70,11 +69,7 @@ func (e *Extractor) Extract(ctx context.Context, channelID, userID, teamID, pers
} }
// Check extraction log for last processed message // Check extraction log for last processed message
var lastMessageID string lastMessageID, _ := e.stores.Memories.GetLastExtractionMessageID(ctx, channelID, userID)
row := database.DB.QueryRowContext(ctx,
database.Q(`SELECT last_message_id FROM memory_extraction_log WHERE channel_id = $1 AND user_id = $2`),
channelID, userID)
row.Scan(&lastMessageID) // ignore error — may not exist yet
// Load recent messages for this channel // Load recent messages for this channel
messages, err := e.stores.Messages.ListForChannel(ctx, channelID, store.ListOptions{Limit: 200}) messages, err := e.stores.Messages.ListForChannel(ctx, channelID, store.ListOptions{Limit: 200})
@@ -197,25 +192,7 @@ func (e *Extractor) Extract(ctx context.Context, channelID, userID, teamID, pers
// Update extraction log // Update extraction log
latestID := newMessages[len(newMessages)-1].ID latestID := newMessages[len(newMessages)-1].ID
if database.IsSQLite() { err = e.stores.Memories.UpsertExtractionLog(ctx, channelID, userID, latestID, saved)
_, err = database.DB.ExecContext(ctx, `
INSERT INTO memory_extraction_log (id, channel_id, user_id, last_message_id, memory_count)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(channel_id, user_id) DO UPDATE SET
last_message_id = excluded.last_message_id,
extracted_at = datetime('now'),
memory_count = memory_extraction_log.memory_count + excluded.memory_count
`, store.NewID(), channelID, userID, latestID, saved)
} else {
_, err = database.DB.ExecContext(ctx, `
INSERT INTO memory_extraction_log (channel_id, user_id, last_message_id, memory_count)
VALUES ($1, $2, $3, $4)
ON CONFLICT(channel_id, user_id) DO UPDATE SET
last_message_id = EXCLUDED.last_message_id,
extracted_at = now(),
memory_count = memory_extraction_log.memory_count + EXCLUDED.memory_count
`, channelID, userID, latestID, saved)
}
if saved > 0 { if saved > 0 {
log.Printf("✅ memory extraction: channel %s → %d facts", channelID, saved) log.Printf("✅ memory extraction: channel %s → %d facts", channelID, saved)
@@ -241,15 +218,7 @@ func (e *Extractor) embedMemory(ctx context.Context, m *models.Memory, userID st
} }
vecJSON, _ := json.Marshal(result.Vectors[0]) vecJSON, _ := json.Marshal(result.Vectors[0])
if database.IsSQLite() { _ = e.stores.Memories.SetEmbedding(ctx, m.ID, string(vecJSON))
database.DB.ExecContext(ctx,
`UPDATE memories SET embedding = ? WHERE id = ?`,
string(vecJSON), m.ID)
} else {
database.DB.ExecContext(ctx,
`UPDATE memories SET embedding = $1::vector WHERE id = $2`,
string(vecJSON), m.ID)
}
} }
// parseExtractionResponse parses the utility model's JSON response. // parseExtractionResponse parses the utility model's JSON response.

View File

@@ -70,12 +70,7 @@ func AuthOrSession(cfg *config.Config, stores store.Stores, cache *UserStatusCac
} }
// Verify channel exists, is workflow type, and allows anonymous // Verify channel exists, is workflow type, and allows anonymous
var chType string chType, allowAnon, err := stores.Channels.GetTypeAndAllowAnonymous(c.Request.Context(), channelID)
var allowAnon bool
err := database.DB.QueryRowContext(c.Request.Context(),
database.Q(`SELECT type, allow_anonymous FROM channels WHERE id = $1`),
channelID,
).Scan(&chType, &allowAnon)
if err != nil { if err != nil {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "authentication required"}) c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "authentication required"})
return return

View File

@@ -5,16 +5,14 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"git.gobha.me/xcaliber/chat-switchboard/database" "git.gobha.me/xcaliber/chat-switchboard/store"
) )
// RequireTeamAdmin returns middleware that restricts access to users who are // RequireTeamAdmin returns middleware that restricts access to team admins.
// admins of the team identified by :teamId in the URL path. // System admins are always allowed through.
// System admins (role=admin) are always allowed through. // v0.29.0: accepts TeamStore instead of using database.DB directly.
// Must be used after Auth() middleware. func RequireTeamAdmin(teams store.TeamStore) gin.HandlerFunc {
func RequireTeamAdmin() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
// System admins bypass team check
role, _ := c.Get("role") role, _ := c.Get("role")
if role == "admin" { if role == "admin" {
c.Next() c.Next()
@@ -30,13 +28,8 @@ func RequireTeamAdmin() gin.HandlerFunc {
return return
} }
var teamRole string isAdmin, err := teams.IsTeamAdmin(c.Request.Context(), teamID, userID.(string))
err := database.DB.QueryRow(` if err != nil || !isAdmin {
SELECT role FROM team_members
WHERE team_id = $1 AND user_id = $2
`, teamID, userID).Scan(&teamRole)
if err != nil || teamRole != "admin" {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{ c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
"error": "team admin access required", "error": "team admin access required",
}) })
@@ -50,7 +43,8 @@ func RequireTeamAdmin() gin.HandlerFunc {
// RequireTeamMember returns middleware that restricts access to users who // RequireTeamMember returns middleware that restricts access to users who
// belong to the team identified by :teamId (any role). // belong to the team identified by :teamId (any role).
// System admins are always allowed through. // System admins are always allowed through.
func RequireTeamMember() gin.HandlerFunc { // v0.29.0: accepts TeamStore instead of using database.DB directly.
func RequireTeamMember(teams store.TeamStore) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
role, _ := c.Get("role") role, _ := c.Get("role")
if role == "admin" { if role == "admin" {
@@ -67,12 +61,8 @@ func RequireTeamMember() gin.HandlerFunc {
return return
} }
var exists bool isMember, _ := teams.IsMember(c.Request.Context(), teamID, userID.(string))
database.DB.QueryRow(` if !isMember {
SELECT EXISTS(SELECT 1 FROM team_members WHERE team_id = $1 AND user_id = $2)
`, teamID, userID).Scan(&exists)
if !exists {
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{ c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
"error": "team membership required", "error": "team membership required",
}) })

View File

@@ -0,0 +1,58 @@
package models
import "time"
// ── Package Status Constants ─────────────────
const (
// PackageStatusActive means the package is running normally.
PackageStatusActive = "active"
// PackageStatusPendingReview means the package declared permissions
// that require admin approval before activation.
PackageStatusPendingReview = "pending_review"
// PackageStatusSuspended means an admin has suspended the package.
PackageStatusSuspended = "suspended"
)
// ValidPackageStatuses is the set of valid package status values.
var ValidPackageStatuses = map[string]bool{
PackageStatusActive: true,
PackageStatusPendingReview: true,
PackageStatusSuspended: true,
}
// ── Extension Permission Constants ───────────
const (
ExtPermSecretsRead = "secrets.read"
ExtPermNotificationsSend = "notifications.send"
ExtPermFiltersPreCompletion = "filters.pre_completion"
ExtPermDBRead = "db.read"
ExtPermDBWrite = "db.write"
ExtPermAPIHTTP = "api.http"
)
// ValidExtensionPermissions is the set of recognized permission keys.
var ValidExtensionPermissions = map[string]bool{
ExtPermSecretsRead: true,
ExtPermNotificationsSend: true,
ExtPermFiltersPreCompletion: true,
ExtPermDBRead: true,
ExtPermDBWrite: true,
ExtPermAPIHTTP: true,
}
// ── Extension Permission Model ───────────────
// ExtensionPermission represents a declared/granted capability for a package.
type ExtensionPermission struct {
ID string `json:"id" db:"id"`
PackageID string `json:"package_id" db:"package_id"`
Permission string `json:"permission" db:"permission"`
Granted bool `json:"granted" db:"granted"`
GrantedBy *string `json:"granted_by,omitempty" db:"granted_by"`
GrantedAt *time.Time `json:"granted_at,omitempty" db:"granted_at"`
CreatedAt time.Time `json:"created_at" db:"created_at"`
}

136
server/sandbox/modules.go Normal file
View File

@@ -0,0 +1,136 @@
// Package sandbox — modules.go
//
// v0.29.0 CS3: Module factories for Starlark sandbox.
// Each factory returns a starlarkstruct.Module that the runner
// injects into the script namespace based on granted permissions.
package sandbox
import (
"context"
"fmt"
"go.starlark.net/starlark"
"go.starlark.net/starlarkstruct"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ─── Secrets Module ──────────────────────────
//
// Requires permission: secrets.read
// Starlark API:
// val = secrets.get("api_key") → string or None
// all = secrets.list() → list of key names
//
// Secrets are stored in GlobalConfig under key "ext_secrets:{packageID}"
// as a JSON map: {"api_key": "sk-...", "webhook_token": "tok-..."}
// Admin sets them via PUT /admin/extensions/:id/secrets.
// BuildSecretsModule creates the "secrets" module scoped to a package.
func BuildSecretsModule(ctx context.Context, stores store.Stores, packageID string) *starlarkstruct.Module {
// Pre-load secrets for this package (single query, not per-call)
secretMap := loadExtensionSecrets(ctx, stores, packageID)
return MakeModule("secrets", starlark.StringDict{
"get": starlark.NewBuiltin("secrets.get", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
var key string
if err := starlark.UnpackPositionalArgs(b.Name(), args, kwargs, 1, &key); err != nil {
return nil, err
}
if val, ok := secretMap[key]; ok {
return starlark.String(val), nil
}
return starlark.None, nil
}),
"list": starlark.NewBuiltin("secrets.list", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
keys := make([]starlark.Value, 0, len(secretMap))
for k := range secretMap {
keys = append(keys, starlark.String(k))
}
return starlark.NewList(keys), nil
}),
})
}
func loadExtensionSecrets(ctx context.Context, stores store.Stores, packageID string) map[string]string {
result := make(map[string]string)
configKey := "ext_secrets:" + packageID
configVal, err := stores.GlobalConfig.Get(ctx, configKey)
if err != nil || configVal == nil {
return result
}
for k, v := range configVal {
if s, ok := v.(string); ok {
result[k] = s
}
}
return result
}
// ─── Notifications Module ────────────────────
//
// Requires permission: notifications.send
// Starlark API:
// notifications.send(user_id, title, body="", type="extension.notify")
//
// Creates an in-app notification via the notification store.
// The extension cannot send email or bypass user preferences —
// that's handled by the notification service layer.
// NotificationSender is the interface the notifications module needs.
// Matches notifications.Service.Notify without importing the package.
type NotificationSender interface {
Notify(ctx context.Context, n *models.Notification) error
}
// BuildNotificationsModule creates the "notifications" module scoped to a package.
func BuildNotificationsModule(ctx context.Context, sender NotificationSender, packageID string) *starlarkstruct.Module {
return MakeModule("notifications", starlark.StringDict{
"send": starlark.NewBuiltin("notifications.send", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
var userID, title string
var body string
var notifType string = "extension.notify"
if err := starlark.UnpackArgs(b.Name(), args, kwargs,
"user_id", &userID,
"title", &title,
"body?", &body,
"type?", &notifType,
); err != nil {
return nil, err
}
if userID == "" {
return nil, fmt.Errorf("notifications.send: user_id is required")
}
if title == "" {
return nil, fmt.Errorf("notifications.send: title is required")
}
n := &models.Notification{
UserID: userID,
Type: notifType,
Title: title,
Body: body,
ResourceType: "extension",
ResourceID: packageID,
}
if err := sender.Notify(ctx, n); err != nil {
return nil, fmt.Errorf("notifications.send failed: %w", err)
}
return starlark.True, nil
}),
})
}

View File

@@ -0,0 +1,171 @@
package sandbox
import (
"context"
"testing"
"go.starlark.net/starlark"
)
func TestBuildSecretsModule_Get(t *testing.T) {
mod := MakeModule("secrets", starlark.StringDict{
"get": starlark.NewBuiltin("secrets.get", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
var key string
if err := starlark.UnpackPositionalArgs(b.Name(), args, kwargs, 1, &key); err != nil {
return nil, err
}
if key == "api_key" {
return starlark.String("sk-test123"), nil
}
return starlark.None, nil
}),
})
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "test.star", `
val = secrets.get("api_key")
missing = secrets.get("nope")
`, map[string]starlark.Value{"secrets": mod})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
val := res.Globals["val"]
if val.(starlark.String).GoString() != "sk-test123" {
t.Fatalf("expected 'sk-test123', got %s", val.String())
}
missing := res.Globals["missing"]
if missing != starlark.None {
t.Fatalf("expected None, got %s", missing.String())
}
}
func TestBuildNotificationsModule_Send(t *testing.T) {
var sentTitle string
var sentUserID string
sendFn := starlark.NewBuiltin("notifications.send", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
var userID, title string
if err := starlark.UnpackArgs(b.Name(), args, kwargs,
"user_id", &userID, "title", &title); err != nil {
return nil, err
}
sentUserID = userID
sentTitle = title
return starlark.True, nil
})
mod := MakeModule("notifications", starlark.StringDict{
"send": sendFn,
})
sb := New(DefaultConfig())
_, err := sb.Exec(context.Background(), "test.star", `
result = notifications.send(user_id="u-1", title="Hello World")
`, map[string]starlark.Value{"notifications": mod})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if sentUserID != "u-1" {
t.Fatalf("expected user_id 'u-1', got %q", sentUserID)
}
if sentTitle != "Hello World" {
t.Fatalf("expected title 'Hello World', got %q", sentTitle)
}
}
func TestModuleIsolation(t *testing.T) {
// Script should not be able to access modules it wasn't given
sb := New(DefaultConfig())
_, err := sb.Exec(context.Background(), "test.star", `
val = secrets.get("key")
`, nil) // no modules injected
if err == nil {
t.Fatal("expected error when accessing undefined module")
}
}
func TestMakeModuleName(t *testing.T) {
mod := MakeModule("test_mod", starlark.StringDict{
"version": starlark.String("1.0"),
})
if mod.Name != "test_mod" {
t.Fatalf("expected 'test_mod', got %q", mod.Name)
}
}
func TestEntryPointCall(t *testing.T) {
sb := New(DefaultConfig())
// Simulate a package script with on_pre_completion entry point
res, err := sb.Exec(context.Background(), "ext.star", `
def on_pre_completion(ctx):
channel = ctx["channel_id"]
return [{"role": "system", "content": "context for " + channel}]
`, nil)
if err != nil {
t.Fatalf("exec error: %v", err)
}
fn := res.Globals["on_pre_completion"].(starlark.Callable)
ctxDict := starlark.NewDict(1)
ctxDict.SetKey(starlark.String("channel_id"), starlark.String("ch-abc"))
val, _, err := sb.Call(context.Background(), fn, starlark.Tuple{ctxDict}, nil)
if err != nil {
t.Fatalf("call error: %v", err)
}
list, ok := val.(*starlark.List)
if !ok {
t.Fatalf("expected list, got %s", val.Type())
}
if list.Len() != 1 {
t.Fatalf("expected 1 message, got %d", list.Len())
}
item := list.Index(0)
dict, ok := item.(*starlark.Dict)
if !ok {
t.Fatalf("expected dict, got %s", item.Type())
}
contentVal, _, _ := dict.Get(starlark.String("content"))
content, _ := starlark.AsString(contentVal)
if content != "context for ch-abc" {
t.Fatalf("expected 'context for ch-abc', got %q", content)
}
}
func TestEntryPointNoneReturn(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "ext.star", `
def on_pre_completion(ctx):
return None
`, nil)
if err != nil {
t.Fatalf("exec error: %v", err)
}
fn := res.Globals["on_pre_completion"].(starlark.Callable)
ctxDict := starlark.NewDict(0)
val, _, err := sb.Call(context.Background(), fn, starlark.Tuple{ctxDict}, nil)
if err != nil {
t.Fatalf("call error: %v", err)
}
if val != starlark.None {
t.Fatalf("expected None, got %s", val.Type())
}
}

131
server/sandbox/runner.go Normal file
View File

@@ -0,0 +1,131 @@
// Package sandbox — runner.go
//
// v0.29.0 CS3: Runner loads a package's Starlark script, assembles
// the module set based on granted permissions, and executes it.
//
// The runner is the bridge between the package/permission system
// and the sandboxed Starlark interpreter.
package sandbox
import (
"context"
"fmt"
"log"
"go.starlark.net/starlark"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// Runner executes Starlark package scripts with permission-gated modules.
type Runner struct {
sandbox *Sandbox
stores store.Stores
notifier NotificationSender // nil = notifications module unavailable
}
// NewRunner creates a runner with the given sandbox and dependencies.
func NewRunner(sb *Sandbox, stores store.Stores) *Runner {
return &Runner{
sandbox: sb,
stores: stores,
}
}
// SetNotifier attaches the notification sender for the notifications module.
func (r *Runner) SetNotifier(n NotificationSender) {
r.notifier = n
}
// ExecPackage loads a package's script from its manifest and executes it
// with modules gated by the package's granted permissions.
//
// Returns the script's globals (which may contain callable entry points
// like on_pre_completion, on_run, etc.) and any captured print output.
func (r *Runner) ExecPackage(ctx context.Context, pkg *store.PackageRegistration) (*Result, error) {
// Only active starlark packages can execute
if pkg.Status != models.PackageStatusActive {
return nil, fmt.Errorf("package %q is %s, not active", pkg.ID, pkg.Status)
}
if pkg.Tier != models.ExtTierStarlark {
return nil, fmt.Errorf("package %q is tier %s, not starlark", pkg.ID, pkg.Tier)
}
// Extract script from manifest
script, ok := pkg.Manifest["_starlark_script"].(string)
if !ok || script == "" {
return nil, fmt.Errorf("package %q has no _starlark_script in manifest", pkg.ID)
}
// Build modules based on granted permissions
modules, err := r.buildModules(ctx, pkg.ID)
if err != nil {
return nil, fmt.Errorf("failed to build modules for %q: %w", pkg.ID, err)
}
log.Printf(" 🔧 runner: exec %s (%d modules granted)", pkg.ID, len(modules))
return r.sandbox.Exec(ctx, pkg.ID+".star", script, modules)
}
// CallEntryPoint executes a package script and calls a named function.
// This is the standard pattern for event-driven extensions:
// 1. Exec the script (defines functions)
// 2. Find the named entry point in globals
// 3. Call it with the provided arguments
//
// Returns the function's return value and captured output.
func (r *Runner) CallEntryPoint(ctx context.Context, pkg *store.PackageRegistration, entryPoint string, args starlark.Tuple, kwargs []starlark.Tuple) (starlark.Value, string, error) {
result, err := r.ExecPackage(ctx, pkg)
if err != nil {
return nil, "", err
}
fn, ok := result.Globals[entryPoint]
if !ok {
return nil, result.Output, fmt.Errorf("package %q has no %s function", pkg.ID, entryPoint)
}
callable, ok := fn.(starlark.Callable)
if !ok {
return nil, result.Output, fmt.Errorf("package %q: %s is not callable", pkg.ID, entryPoint)
}
val, callOutput, err := r.sandbox.Call(ctx, callable, args, kwargs)
return val, result.Output + callOutput, err
}
// buildModules assembles the module map based on granted permissions.
func (r *Runner) buildModules(ctx context.Context, packageID string) (map[string]starlark.Value, error) {
if r.stores.ExtPermissions == nil {
return nil, nil
}
granted, err := r.stores.ExtPermissions.GrantedForPackage(ctx, packageID)
if err != nil {
return nil, err
}
modules := make(map[string]starlark.Value)
for _, perm := range granted {
switch perm {
case models.ExtPermSecretsRead:
modules["secrets"] = BuildSecretsModule(ctx, r.stores, packageID)
case models.ExtPermNotificationsSend:
if r.notifier != nil {
modules["notifications"] = BuildNotificationsModule(ctx, r.notifier, packageID)
}
// Future permissions (CS4+):
// case models.ExtPermDBRead, models.ExtPermDBWrite:
// modules["db"] = BuildDBModule(...)
// case models.ExtPermAPIHTTP:
// modules["http"] = BuildHTTPModule(...)
}
}
return modules, nil
}

229
server/sandbox/sandbox.go Normal file
View File

@@ -0,0 +1,229 @@
// Package sandbox — sandboxed Starlark interpreter.
//
// v0.29.0 CS1: Provides execution with timeout, step limits,
// captured output, and injectable predeclared modules.
//
// The sandbox is intentionally restrictive:
// - No file I/O
// - No network access
// - No os/env access
// - No load() statements (modules injected via predeclared)
// - Step limit prevents infinite loops
// - Context-based timeout prevents runaway wall-clock
//
// Extensions interact with the platform through injected Go builtins
// (e.g., secrets.get, notifications.send) — never through Starlark's
// own I/O. This is the eval loop; modules are added in CS3.
package sandbox
import (
"context"
"fmt"
"strings"
"sync"
"go.starlark.net/starlark"
"go.starlark.net/starlarkstruct"
)
// ─── Configuration ───────────────────────────
// Config controls sandbox resource limits.
type Config struct {
// MaxSteps limits total bytecode operations. 0 = unlimited.
// Default: 1_000_000 (~1M ops, covers most scripts in <100ms).
MaxSteps uint64
// Name identifies this sandbox in logs and error messages.
Name string
}
// DefaultConfig returns production defaults.
func DefaultConfig() Config {
return Config{
MaxSteps: 1_000_000,
Name: "extension",
}
}
// ─── Result ──────────────────────────────────
// Result holds the output of a sandbox execution.
type Result struct {
// Globals contains the top-level names defined by the script.
Globals map[string]starlark.Value
// Output captures all print() calls made during execution.
Output string
// Steps is the number of bytecode operations executed.
Steps uint64
}
// ─── Sandbox ─────────────────────────────────
// Sandbox executes Starlark scripts with resource limits and
// injected modules. Safe for concurrent use — each Exec call
// gets its own thread.
type Sandbox struct {
config Config
}
// New creates a sandbox with the given configuration.
// Set MaxSteps to 0 for unlimited (not recommended in production).
func New(cfg Config) *Sandbox {
if cfg.Name == "" {
cfg.Name = "extension"
}
return &Sandbox{config: cfg}
}
// Exec executes a Starlark script within the sandbox.
//
// Parameters:
// - ctx: context for timeout/cancellation
// - filename: used in error messages and stack traces
// - source: Starlark source code
// - modules: predeclared modules injected into the script's namespace
// (e.g., {"secrets": secretsModule, "notifications": notifModule})
//
// The script cannot use load() — all dependencies must be provided
// via the modules parameter.
func (s *Sandbox) Exec(ctx context.Context, filename, source string, modules map[string]starlark.Value) (*Result, error) {
var output strings.Builder
var outputMu sync.Mutex
predeclared := make(starlark.StringDict, len(modules))
for k, v := range modules {
predeclared[k] = v
}
thread := &starlark.Thread{
Name: s.config.Name,
Print: func(_ *starlark.Thread, msg string) {
outputMu.Lock()
defer outputMu.Unlock()
output.WriteString(msg)
output.WriteByte('\n')
},
Load: func(_ *starlark.Thread, module string) (starlark.StringDict, error) {
return nil, fmt.Errorf("load() is disabled in sandbox (requested %q)", module)
},
}
if s.config.MaxSteps > 0 {
thread.SetMaxExecutionSteps(s.config.MaxSteps)
}
// Wire context cancellation into the thread.
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
thread.Cancel(ctx.Err().Error())
case <-done:
}
}()
globals, err := starlark.ExecFile(thread, filename, source, predeclared)
if err != nil {
return nil, wrapError(err)
}
result := &Result{
Globals: make(map[string]starlark.Value, len(globals)),
Output: output.String(),
Steps: thread.ExecutionSteps(),
}
for k, v := range globals {
result.Globals[k] = v
}
return result, nil
}
// Call invokes a Starlark function within the sandbox.
// Used by the task executor and filter runner to call specific
// entry points (e.g., "on_pre_completion", "on_run").
func (s *Sandbox) Call(ctx context.Context, fn starlark.Callable, args starlark.Tuple, kwargs []starlark.Tuple) (starlark.Value, string, error) {
var output strings.Builder
var outputMu sync.Mutex
thread := &starlark.Thread{
Name: s.config.Name + "/call",
Print: func(_ *starlark.Thread, msg string) {
outputMu.Lock()
defer outputMu.Unlock()
output.WriteString(msg)
output.WriteByte('\n')
},
Load: func(_ *starlark.Thread, module string) (starlark.StringDict, error) {
return nil, fmt.Errorf("load() is disabled in sandbox (requested %q)", module)
},
}
if s.config.MaxSteps > 0 {
thread.SetMaxExecutionSteps(s.config.MaxSteps)
}
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
thread.Cancel(ctx.Err().Error())
case <-done:
}
}()
val, err := starlark.Call(thread, fn, args, kwargs)
if err != nil {
return nil, output.String(), wrapError(err)
}
return val, output.String(), nil
}
// ─── Module Helper ───────────────────────────
// MakeModule creates a starlarkstruct.Module with the given name and
// members. Standard way to expose Go functions to Starlark.
//
// Usage:
//
// mod := sandbox.MakeModule("secrets", starlark.StringDict{
// "get": starlark.NewBuiltin("secrets.get", secretsGet),
// })
func MakeModule(name string, members starlark.StringDict) *starlarkstruct.Module {
return &starlarkstruct.Module{
Name: name,
Members: members,
}
}
// ─── Error Wrapping ──────────────────────────
// SandboxError wraps a Starlark execution error with context.
type SandboxError struct {
Err error
Backtrace string
}
func (e *SandboxError) Error() string {
if e.Backtrace != "" {
return fmt.Sprintf("%s\n%s", e.Err.Error(), e.Backtrace)
}
return e.Err.Error()
}
func (e *SandboxError) Unwrap() error { return e.Err }
func wrapError(err error) error {
if err == nil {
return nil
}
se := &SandboxError{Err: err}
if evalErr, ok := err.(*starlark.EvalError); ok {
se.Backtrace = evalErr.Backtrace()
}
return se
}

View File

@@ -0,0 +1,306 @@
package sandbox
import (
"context"
"strings"
"testing"
"time"
"go.starlark.net/starlark"
)
func TestExecBasic(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "test.star", `
x = 1 + 2
name = "hello"
`, nil)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if v, ok := res.Globals["x"]; !ok {
t.Fatal("missing global 'x'")
} else if v.String() != "3" {
t.Fatalf("expected x=3, got %s", v.String())
}
if v, ok := res.Globals["name"]; !ok {
t.Fatal("missing global 'name'")
} else if v.(starlark.String).GoString() != "hello" {
t.Fatalf("expected name='hello', got %s", v.String())
}
}
func TestExecPrintCapture(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "test.star", `
print("line one")
print("line two")
`, nil)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
expected := "line one\nline two\n"
if res.Output != expected {
t.Fatalf("expected output %q, got %q", expected, res.Output)
}
}
func TestExecStepsTracked(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "test.star", `
def work():
x = 0
for i in range(100):
x += i
return x
result = work()
`, nil)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if res.Steps == 0 {
t.Fatal("expected steps > 0")
}
}
func TestExecStepLimitExceeded(t *testing.T) {
sb := New(Config{MaxSteps: 100, Name: "limited"})
_, err := sb.Exec(context.Background(), "test.star", `
def work():
x = 0
for i in range(10000):
x += i
return x
result = work()
`, nil)
if err == nil {
t.Fatal("expected step limit error")
}
// Error message varies: "Starlark computation cancelled: exceeded ..."
// or "context canceled" — just verify it failed.
}
func TestExecContextTimeout(t *testing.T) {
// Use step limit of 0 (unlimited) so only context timeout triggers
sb := New(Config{MaxSteps: 0, Name: "timeout"})
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
defer cancel()
_, err := sb.Exec(ctx, "test.star", `
def work():
x = 0
for i in range(100000000):
x += i
return x
result = work()
`, nil)
if err == nil {
t.Fatal("expected timeout error")
}
}
func TestExecLoadDisabled(t *testing.T) {
sb := New(DefaultConfig())
_, err := sb.Exec(context.Background(), "test.star", `
load("other.star", "foo")
`, nil)
if err == nil {
t.Fatal("expected load error")
}
if !strings.Contains(err.Error(), "disabled") {
t.Fatalf("expected 'disabled' in error, got: %v", err)
}
}
func TestExecModuleInjection(t *testing.T) {
getFn := starlark.NewBuiltin("test.get", func(
thread *starlark.Thread, b *starlark.Builtin,
args starlark.Tuple, kwargs []starlark.Tuple,
) (starlark.Value, error) {
return starlark.String("secret_value"), nil
})
mod := MakeModule("test", starlark.StringDict{
"get": getFn,
})
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "test.star", `
result = test.get()
`, map[string]starlark.Value{"test": mod})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
v, ok := res.Globals["result"]
if !ok {
t.Fatal("missing global 'result'")
}
if v.(starlark.String).GoString() != "secret_value" {
t.Fatalf("expected 'secret_value', got %s", v.String())
}
}
func TestExecSyntaxError(t *testing.T) {
sb := New(DefaultConfig())
_, err := sb.Exec(context.Background(), "bad.star", `
def broken(
`, nil)
if err == nil {
t.Fatal("expected syntax error")
}
se, ok := err.(*SandboxError)
if !ok {
t.Fatalf("expected *SandboxError, got %T", err)
}
_ = se // syntax errors may not have backtrace
}
func TestExecRuntimeError(t *testing.T) {
sb := New(DefaultConfig())
_, err := sb.Exec(context.Background(), "runtime.star", `
x = 1 / 0
`, nil)
if err == nil {
t.Fatal("expected runtime error")
}
se, ok := err.(*SandboxError)
if !ok {
t.Fatalf("expected *SandboxError, got %T", err)
}
if se.Backtrace == "" {
t.Fatal("expected backtrace for runtime error")
}
}
func TestCallFunction(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "funcs.star", `
def greet(name):
return "hello, " + name
`, nil)
if err != nil {
t.Fatalf("exec error: %v", err)
}
fn := res.Globals["greet"].(starlark.Callable)
val, _, err := sb.Call(context.Background(), fn,
starlark.Tuple{starlark.String("world")}, nil)
if err != nil {
t.Fatalf("call error: %v", err)
}
if val.(starlark.String).GoString() != "hello, world" {
t.Fatalf("expected 'hello, world', got %s", val.String())
}
}
func TestCallWithPrint(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "funcs.star", `
def run():
print("running")
return 42
`, nil)
if err != nil {
t.Fatalf("exec error: %v", err)
}
fn := res.Globals["run"].(starlark.Callable)
val, output, err := sb.Call(context.Background(), fn, nil, nil)
if err != nil {
t.Fatalf("call error: %v", err)
}
if val.String() != "42" {
t.Fatalf("expected 42, got %s", val.String())
}
if !strings.Contains(output, "running") {
t.Fatalf("expected 'running' in output, got %q", output)
}
}
func TestCallRuntimeError(t *testing.T) {
sb := New(DefaultConfig())
res, err := sb.Exec(context.Background(), "funcs.star", `
def boom():
return 1 / 0
`, nil)
if err != nil {
t.Fatalf("exec error: %v", err)
}
fn := res.Globals["boom"].(starlark.Callable)
_, _, err = sb.Call(context.Background(), fn, nil, nil)
if err == nil {
t.Fatal("expected runtime error from Call")
}
se, ok := err.(*SandboxError)
if !ok {
t.Fatalf("expected *SandboxError, got %T", err)
}
if se.Backtrace == "" {
t.Fatal("expected backtrace")
}
}
func TestNewDefaults(t *testing.T) {
sb := New(Config{})
if sb.config.MaxSteps != 0 {
t.Fatalf("expected MaxSteps=0 (unlimited), got %d", sb.config.MaxSteps)
}
if sb.config.Name != "extension" {
t.Fatalf("expected default Name='extension', got %q", sb.config.Name)
}
}
func TestNewWithDefaultConfig(t *testing.T) {
sb := New(DefaultConfig())
if sb.config.MaxSteps != 1_000_000 {
t.Fatalf("expected MaxSteps=1000000, got %d", sb.config.MaxSteps)
}
}
func TestMakeModule(t *testing.T) {
mod := MakeModule("mymod", starlark.StringDict{
"version": starlark.String("1.0"),
})
if mod.Name != "mymod" {
t.Fatalf("expected name 'mymod', got %q", mod.Name)
}
v, ok := mod.Members["version"]
if !ok {
t.Fatal("missing 'version' member")
}
if v.(starlark.String).GoString() != "1.0" {
t.Fatalf("expected '1.0', got %s", v.String())
}
}
func TestSandboxErrorUnwrap(t *testing.T) {
inner := starlark.String("test").Truth() // just need any value
_ = inner
se := &SandboxError{
Err: context.DeadlineExceeded,
Backtrace: "",
}
if se.Unwrap() != context.DeadlineExceeded {
t.Fatal("Unwrap should return inner error")
}
if se.Error() != "context deadline exceeded" {
t.Fatalf("unexpected error string: %q", se.Error())
}
}
func TestSandboxErrorWithBacktrace(t *testing.T) {
se := &SandboxError{
Err: context.DeadlineExceeded,
Backtrace: " funcs.star:3:5: in boom",
}
if !strings.Contains(se.Error(), "context deadline exceeded") {
t.Fatal("should contain original error")
}
if !strings.Contains(se.Error(), "funcs.star:3:5") {
t.Fatal("should contain backtrace")
}
}

View File

@@ -22,6 +22,7 @@ import (
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/notifications" "git.gobha.me/xcaliber/chat-switchboard/notifications"
"git.gobha.me/xcaliber/chat-switchboard/providers" "git.gobha.me/xcaliber/chat-switchboard/providers"
"git.gobha.me/xcaliber/chat-switchboard/sandbox"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
"git.gobha.me/xcaliber/chat-switchboard/taskutil" "git.gobha.me/xcaliber/chat-switchboard/taskutil"
"git.gobha.me/xcaliber/chat-switchboard/tools" "git.gobha.me/xcaliber/chat-switchboard/tools"
@@ -34,6 +35,7 @@ type Executor struct {
vault *crypto.KeyResolver vault *crypto.KeyResolver
hub *events.Hub hub *events.Hub
health handlers.HealthRecorder health handlers.HealthRecorder
runner *sandbox.Runner // v0.29.0: Starlark task execution
} }
// NewExecutor creates a task executor. All fields are optional except stores. // NewExecutor creates a task executor. All fields are optional except stores.
@@ -46,6 +48,11 @@ func NewExecutor(stores store.Stores, vault *crypto.KeyResolver, hub *events.Hub
} }
} }
// SetRunner attaches the Starlark sandbox runner for starlark task execution.
func (e *Executor) SetRunner(r *sandbox.Runner) {
e.runner = r
}
// Execute runs a single task to completion. // Execute runs a single task to completion.
// Called from the scheduler's execute() goroutine. // Called from the scheduler's execute() goroutine.
func (e *Executor) Execute(ctx context.Context, task models.Task, run *models.TaskRun, channelID string) { func (e *Executor) Execute(ctx context.Context, task models.Task, run *models.TaskRun, channelID string) {
@@ -63,12 +70,18 @@ func (e *Executor) Execute(ctx context.Context, task models.Task, run *models.Ta
return return
} }
// v0.29.0: Starlark tasks run a sandboxed extension script.
if task.TaskType == "starlark" {
e.executeStarlark(ctx, task, run, channelID, startTime)
return
}
// ── 1. Resolve provider ──────────────────── // ── 1. Resolve provider ────────────────────
providerConfigID := "" providerConfigID := ""
if task.ProviderConfigID != nil { if task.ProviderConfigID != nil {
providerConfigID = *task.ProviderConfigID providerConfigID = *task.ProviderConfigID
} }
res, err := handlers.ResolveProviderConfig(e.vault, task.OwnerID, channelID, providerConfigID, task.ModelID) res, err := handlers.ResolveProviderConfig(e.stores, e.vault, task.OwnerID, channelID, providerConfigID, task.ModelID)
if err != nil { if err != nil {
e.failRun(ctx, task, run, "provider resolution failed: "+err.Error()) e.failRun(ctx, task, run, "provider resolution failed: "+err.Error())
return return
@@ -265,6 +278,83 @@ func (e *Executor) Execute(ctx context.Context, task models.Task, run *models.Ta
} }
} }
// executeStarlark runs a sandboxed Starlark extension script.
// The task's system_function field holds the package ID.
// The script's on_run(ctx) entry point is called with task context.
func (e *Executor) executeStarlark(ctx context.Context, task models.Task, run *models.TaskRun, channelID string, startTime time.Time) {
if e.runner == nil {
e.failRun(ctx, task, run, "starlark runner not configured")
return
}
packageID := task.SystemFunction
if packageID == "" {
e.failRun(ctx, task, run, "starlark task missing package_id (system_function field)")
return
}
// Load the package
pkg, err := e.stores.Packages.Get(ctx, packageID)
if err != nil || pkg == nil {
e.failRun(ctx, task, run, "package not found: "+packageID)
return
}
// Run the script and call on_run entry point
// Build a context dict with task info
val, output, err := e.runner.CallEntryPoint(ctx, pkg, "on_run", nil, nil)
wallClock := int(time.Since(startTime).Seconds())
status := "completed"
errMsg := ""
result := ""
if err != nil {
status = "failed"
errMsg = err.Error()
} else if val != nil {
result = val.String()
}
// Append print output to result
if output != "" {
if result != "" {
result = result + "\n---\n" + output
} else {
result = output
}
}
// Persist to channel if output_mode == "channel"
if status == "completed" && task.OutputMode == "channel" && e.stores.Messages != nil && result != "" {
_ = e.stores.Messages.Create(ctx, &models.Message{
ChannelID: channelID,
Role: "system",
Content: "Starlark task output:\n```\n" + result + "\n```",
})
}
_ = e.stores.Tasks.UpdateRun(ctx, run.ID, status, 0, 0, wallClock, errMsg)
_ = e.stores.Tasks.IncrementRunCount(ctx, task.ID)
log.Printf("[executor] Starlark task %s (%s → %s) → %s (wall=%ds)", task.ID, task.Name, packageID, status, wallClock)
e.notifyOwner(ctx, task, status, errMsg)
if task.WebhookURL != "" {
go webhook.Deliver(task.WebhookURL, task.WebhookSecret, webhook.Payload{
TaskID: task.ID,
RunID: run.ID,
TaskName: task.Name,
ChannelID: channelID,
Status: status,
CompletedAt: time.Now().UTC(),
Output: result,
Error: errMsg,
})
}
}
// executeSystem runs a built-in Go function from the system registry. // executeSystem runs a built-in Go function from the system registry.
// No LLM, no provider resolution, no channel needed. // No LLM, no provider resolution, no channel needed.
func (e *Executor) executeSystem(ctx context.Context, task models.Task, run *models.TaskRun, startTime time.Time) { func (e *Executor) executeSystem(ctx context.Context, task models.Task, run *models.TaskRun, startTime time.Time) {

View File

@@ -1,9 +1,7 @@
// Package scheduler — system_builtins.go // Package scheduler — system_builtins.go
// //
// v0.28.6: Built-in system functions registered at startup. // v0.28.6: Built-in system functions registered at startup.
// These mirror the background goroutines in main.go but run as visible, // v0.29.0: Raw SQL replaced with store methods.
// scheduled, auditable tasks. The goroutines remain as fallback until
// system tasks are validated in production.
package scheduler package scheduler
import ( import (
@@ -12,13 +10,11 @@ import (
"log" "log"
"time" "time"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
"git.gobha.me/xcaliber/chat-switchboard/taskutil" "git.gobha.me/xcaliber/chat-switchboard/taskutil"
) )
// RegisterBuiltins registers all built-in system functions. // RegisterBuiltins registers all built-in system functions.
// Called once from main.go at startup.
func RegisterBuiltins() { func RegisterBuiltins() {
taskutil.RegisterSystemFunc("session_cleanup", taskutil.RegisterSystemFunc("session_cleanup",
"Remove expired anonymous visitor sessions", "Remove expired anonymous visitor sessions",
@@ -43,7 +39,6 @@ func sessionCleanup(ctx context.Context, stores store.Stores) (string, error) {
if stores.Sessions == nil { if stores.Sessions == nil {
return "skipped: session store not available", nil return "skipped: session store not available", nil
} }
// Default: expire sessions older than 7 days
cutoff := time.Now().UTC().AddDate(0, 0, -7) cutoff := time.Now().UTC().AddDate(0, 0, -7)
n, err := stores.Sessions.DeleteExpired(ctx, cutoff) n, err := stores.Sessions.DeleteExpired(ctx, cutoff)
if err != nil { if err != nil {
@@ -59,19 +54,11 @@ func sessionCleanup(ctx context.Context, stores store.Stores) (string, error) {
// ── staleness_check ───────────────────────── // ── staleness_check ─────────────────────────
func stalenessCheck(ctx context.Context, stores store.Stores) (string, error) { func stalenessCheck(ctx context.Context, stores store.Stores) (string, error) {
// Default: 48 hours idle → stale
cutoff := time.Now().UTC().Add(-48 * time.Hour) cutoff := time.Now().UTC().Add(-48 * time.Hour)
res, err := database.DB.ExecContext(ctx, database.Q(` n, err := stores.Channels.MarkStaleWorkflows(ctx, cutoff)
UPDATE channels
SET workflow_status = 'stale'
WHERE type = 'workflow'
AND workflow_status = 'active'
AND last_activity_at < $1
`), cutoff)
if err != nil { if err != nil {
return "", fmt.Errorf("staleness sweep failed: %w", err) return "", fmt.Errorf("staleness sweep failed: %w", err)
} }
n, _ := res.RowsAffected()
msg := fmt.Sprintf("marked %d idle workflow instances as stale (cutoff: %s)", n, cutoff.Format(time.RFC3339)) msg := fmt.Sprintf("marked %d idle workflow instances as stale (cutoff: %s)", n, cutoff.Format(time.RFC3339))
if n > 0 { if n > 0 {
log.Printf("[system_task] %s", msg) log.Printf("[system_task] %s", msg)
@@ -82,28 +69,10 @@ func stalenessCheck(ctx context.Context, stores store.Stores) (string, error) {
// ── retention_sweep ───────────────────────── // ── retention_sweep ─────────────────────────
func retentionSweep(ctx context.Context, stores store.Stores) (string, error) { func retentionSweep(ctx context.Context, stores store.Stores) (string, error) {
if database.IsSQLite() { n, err := stores.Channels.EnforceWorkflowRetention(ctx)
return "skipped: retention enforcement requires PostgreSQL (JSON operators)", nil
}
cutoff := time.Now().UTC().Add(-48 * time.Hour)
res, err := database.DB.ExecContext(ctx, database.Q(`
DELETE FROM channels
WHERE type = 'workflow'
AND workflow_status IN ('completed', 'archived')
AND workflow_id IS NOT NULL
AND last_activity_at < $1
AND workflow_id IN (
SELECT id FROM workflows
WHERE retention IS NOT NULL
AND retention->>'mode' = 'delete'
AND (retention->>'delete_after_days')::int > 0
AND channels.last_activity_at < now() - ((retention->>'delete_after_days')::int || ' days')::interval
)
`), cutoff)
if err != nil { if err != nil {
return "", fmt.Errorf("retention enforcement failed: %w", err) return "", fmt.Errorf("retention enforcement failed: %w", err)
} }
n, _ := res.RowsAffected()
msg := fmt.Sprintf("deleted %d expired workflow instances (retention policy)", n) msg := fmt.Sprintf("deleted %d expired workflow instances (retention policy)", n)
if n > 0 { if n > 0 {
log.Printf("[system_task] %s", msg) log.Printf("[system_task] %s", msg)
@@ -114,14 +83,14 @@ func retentionSweep(ctx context.Context, stores store.Stores) (string, error) {
// ── health_prune ──────────────────────────── // ── health_prune ────────────────────────────
func healthPrune(ctx context.Context, stores store.Stores) (string, error) { func healthPrune(ctx context.Context, stores store.Stores) (string, error) {
if stores.Health == nil {
return "skipped: health store not available", nil
}
cutoff := time.Now().UTC().Add(-7 * 24 * time.Hour) cutoff := time.Now().UTC().Add(-7 * 24 * time.Hour)
res, err := database.DB.ExecContext(ctx, database.Q(` n, err := stores.Health.Prune(ctx, cutoff)
DELETE FROM health_windows WHERE window_start < $1
`), cutoff)
if err != nil { if err != nil {
return "", fmt.Errorf("health prune failed: %w", err) return "", fmt.Errorf("health prune failed: %w", err)
} }
n, _ := res.RowsAffected()
msg := fmt.Sprintf("pruned %d old health windows (cutoff: %s)", n, cutoff.Format(time.RFC3339)) msg := fmt.Sprintf("pruned %d old health windows (cutoff: %s)", n, cutoff.Format(time.RFC3339))
if n > 0 { if n > 0 {
log.Printf("[system_task] %s", msg) log.Printf("[system_task] %s", msg)

View File

@@ -0,0 +1,36 @@
package store
import (
"context"
"git.gobha.me/xcaliber/chat-switchboard/models"
)
// ExtensionPermissionStore manages declared and granted permissions for packages.
type ExtensionPermissionStore interface {
// DeclareForPackage upserts the declared permissions from a manifest.
// Any permissions in the DB for this package that are NOT in the
// provided list are deleted (manifest is the source of truth).
// Existing grants are preserved on upsert.
DeclareForPackage(ctx context.Context, packageID string, permissions []string) error
// ListForPackage returns all declared permissions for a package.
ListForPackage(ctx context.Context, packageID string) ([]models.ExtensionPermission, error)
// GrantedForPackage returns only granted permissions for a package.
// Used at runtime to determine which modules to inject.
GrantedForPackage(ctx context.Context, packageID string) ([]string, error)
// Grant marks a permission as granted by an admin.
Grant(ctx context.Context, packageID, permission, grantedBy string) error
// Revoke removes a grant (sets granted=false).
Revoke(ctx context.Context, packageID, permission string) error
// GrantAll grants all declared permissions for a package.
GrantAll(ctx context.Context, packageID, grantedBy string) error
// DeleteForPackage removes all permission rows for a package.
// Called when a package is uninstalled.
DeleteForPackage(ctx context.Context, packageID string) error
}

View File

@@ -2,6 +2,7 @@ package store
import ( import (
"context" "context"
"encoding/json"
"errors" "errors"
"time" "time"
@@ -55,6 +56,11 @@ type Stores struct {
Packages PackageStore // v0.28.7: Unified package registry (surfaces + extensions) Packages PackageStore // v0.28.7: Unified package registry (surfaces + extensions)
Workflows WorkflowStore // v0.26.1: Workflow definitions + stages Workflows WorkflowStore // v0.26.1: Workflow definitions + stages
Tasks TaskStore // v0.27.1: Task scheduling + run history Tasks TaskStore // v0.27.1: Task scheduling + run history
Presence PresenceStore // v0.29.0: User online/offline status
PersonaGroups PersonaGroupStore // v0.29.0: Persona group (roster template) CRUD
Folders FolderStore // v0.29.0: Chat folder CRUD
Health HealthStore // v0.29.0: Provider health window management
ExtPermissions ExtensionPermissionStore // v0.29.0: Extension declared/granted capabilities
} }
// ========================================= // =========================================
@@ -75,6 +81,31 @@ type ProviderStore interface {
// Access check // Access check
UserCanAccess(ctx context.Context, userID, configID string) (bool, error) UserCanAccess(ctx context.Context, userID, configID string) (bool, error)
// ── CS4 additions (v0.29.0) ──
// DeletePersonalByOwner removes all personal-scope provider configs for a user.
DeletePersonalByOwner(ctx context.Context, ownerID string) (int64, error)
// ── CS6 additions (v0.29.0) ──
// ListAllForTeam returns all provider configs (including inactive) for a team.
ListAllForTeam(ctx context.Context, teamID string) ([]models.ProviderConfig, error)
// DeleteByIDAndTeam deletes a team-scoped provider by ID and team. Returns rows affected.
DeleteByIDAndTeam(ctx context.Context, id, teamID string) (int64, error)
// ── CS7a additions (v0.29.0) ──
// FindFirstForUser returns the ID of the first active provider config
// accessible to a user (personal first, then global; excludes team).
// Returns sql.ErrNoRows if none found.
FindFirstForUser(ctx context.Context, userID string) (string, error)
// LoadAccessible loads a provider config by ID, verifying the user has access
// (global, personal owned by user, or team the user belongs to).
// Returns sql.ErrNoRows if not found or not accessible.
LoadAccessible(ctx context.Context, configID, userID string) (*models.ProviderConfig, error)
} }
// ========================================= // =========================================
@@ -103,6 +134,41 @@ type CatalogStore interface {
// Delete // Delete
Delete(ctx context.Context, id string) error Delete(ctx context.Context, id string) error
DeleteForProvider(ctx context.Context, providerConfigID string) error DeleteForProvider(ctx context.Context, providerConfigID string) error
// ── Mention resolution (v0.29.0 — moved from handler raw SQL) ──
// FindEnabledByModelID finds an enabled catalog entry by exact model_id
// (case-insensitive). Returns (modelID, providerConfigID, nil).
// Prefers global > team > personal scope providers.
FindEnabledByModelID(ctx context.Context, modelID string) (string, string, error)
// FindEnabledByModelIDPrefix finds by unambiguous model_id prefix.
// Returns (modelID, configID, matchCount, nil).
FindEnabledByModelIDPrefix(ctx context.Context, prefix string) (string, string, int, error)
// ── CS2 additions (v0.29.0) ──
// GetCapabilities returns the raw capabilities JSON for a specific provider+model.
GetCapabilities(ctx context.Context, modelID, configID string) ([]byte, error)
// GetCapabilitiesAny returns capabilities for a model from any provider (most recent sync).
GetCapabilitiesAny(ctx context.Context, modelID string) ([]byte, error)
// ── CS6 additions (v0.29.0) ──
// ListTeamAvailable returns catalog entries with visibility 'enabled' or 'team'
// from active global providers. Used by team model selector.
ListTeamAvailable(ctx context.Context) ([]TeamAvailableModel, error)
}
// TeamAvailableModel is returned by CatalogStore.ListTeamAvailable.
type TeamAvailableModel struct {
ID string
ModelID string
DisplayName *string
Visibility string
Provider string
ProviderName string
} }
// CatalogSyncEntry is the input format from provider FetchModels. // CatalogSyncEntry is the input format from provider FetchModels.
@@ -141,6 +207,25 @@ type PersonaStore interface {
SetKBs(ctx context.Context, personaID string, kbIDs []string, autoSearch map[string]bool) error SetKBs(ctx context.Context, personaID string, kbIDs []string, autoSearch map[string]bool) error
GetKBs(ctx context.Context, personaID string) ([]models.PersonaKB, error) GetKBs(ctx context.Context, personaID string) ([]models.PersonaKB, error)
GetKBIDs(ctx context.Context, personaID string) ([]string, error) GetKBIDs(ctx context.Context, personaID string) ([]string, error)
// ── Mention resolution (v0.29.0 — moved from handler raw SQL) ──
// FindActiveByHandle returns an active persona's ID by exact handle match (case-insensitive).
// Returns "" if not found.
FindActiveByHandle(ctx context.Context, handle string) (string, error)
// FindActiveByHandlePrefix returns a persona ID by unambiguous handle prefix.
// Returns ("", count, nil) if ambiguous or no match.
FindActiveByHandlePrefix(ctx context.Context, prefix string) (string, int, error)
// GetNameByID returns the persona's display name. Returns "" if not found.
GetNameByID(ctx context.Context, id string) (string, error)
// GetNamesByIDs returns a map of persona ID → name for batch resolution.
GetNamesByIDs(ctx context.Context, ids []string) (map[string]string, error)
// GetDisplayInfoByIDs returns name + avatar for batch sender resolution.
GetDisplayInfoByIDs(ctx context.Context, ids []string) (map[string]UserDisplayInfo, error)
} }
// ========================================= // =========================================
@@ -190,6 +275,61 @@ type UserStore interface {
RevokeRefreshToken(ctx context.Context, tokenHash string) error RevokeRefreshToken(ctx context.Context, tokenHash string) error
RevokeAllRefreshTokens(ctx context.Context, userID string) error RevokeAllRefreshTokens(ctx context.Context, userID string) error
CleanExpiredTokens(ctx context.Context) error CleanExpiredTokens(ctx context.Context) error
// ── Mention resolution (v0.29.0 — moved from handler raw SQL) ──
// FindActiveByHandle returns a user ID by exact handle match (case-insensitive).
// excludeUserID prevents self-mention resolution.
FindActiveByHandle(ctx context.Context, handle, excludeUserID string) (string, error)
// FindActiveByHandlePrefix returns a user ID by unambiguous handle prefix.
// Returns ("", count, nil) if ambiguous or no match.
FindActiveByHandlePrefix(ctx context.Context, prefix, excludeUserID string) (string, int, error)
// GetDisplayInfoByIDs returns name + avatar for batch sender resolution.
GetDisplayInfoByIDs(ctx context.Context, ids []string) (map[string]UserDisplayInfo, error)
// ── CS1 additions (v0.29.0) ──
// Exists returns true if a user with the given ID exists and is active.
Exists(ctx context.Context, userID string) (bool, error)
// SearchActive returns users matching a query (username, display_name, handle).
// Excludes the calling user. Max 20 results.
SearchActive(ctx context.Context, excludeUserID, query string) ([]UserSearchResult, error)
// ── CS2 additions (v0.29.0) ──
// CountByRole returns the number of users with a given role.
CountByRole(ctx context.Context, role string) (int, error)
// CountAll returns the total number of users.
CountAll(ctx context.Context) (int, error)
// MergeSettings performs a dialect-aware JSON merge into the user's settings column.
MergeSettings(ctx context.Context, userID string, patch []byte) error
// GetVaultKeys returns the user's vault encryption state.
GetVaultKeys(ctx context.Context, userID string) (vaultSet bool, encUEK, salt, nonce []byte, err error)
// UpdateVaultKeys re-wraps the user's vault keys (password change).
UpdateVaultKeys(ctx context.Context, userID string, encUEK, salt, nonce []byte) error
// ── CS4 additions (v0.29.0) ──
// ClearVaultKeys nulls out vault columns and sets vault_set = false.
ClearVaultKeys(ctx context.Context, userID string) error
// InitVaultKeys stores vault keys and sets vault_set = true (first-time init).
InitVaultKeys(ctx context.Context, userID string, encUEK, salt, nonce []byte) error
}
// UserSearchResult is a lightweight result for user search.
type UserSearchResult struct {
ID string `json:"id"`
Username string `json:"username"`
DisplayName string `json:"display_name"`
Handle string `json:"handle"`
} }
// ========================================= // =========================================
@@ -213,6 +353,37 @@ type TeamStore interface {
GetUserTeamIDs(ctx context.Context, userID string) ([]string, error) GetUserTeamIDs(ctx context.Context, userID string) ([]string, error)
IsTeamAdmin(ctx context.Context, teamID, userID string) (bool, error) IsTeamAdmin(ctx context.Context, teamID, userID string) (bool, error)
IsMember(ctx context.Context, teamID, userID string) (bool, error) IsMember(ctx context.Context, teamID, userID string) (bool, error)
// ── CS1 additions (v0.29.0) ──
// Exists returns true if a team with the given ID exists.
Exists(ctx context.Context, teamID string) (bool, error)
// UpdateMemberRoleByID updates a member's role using the member row ID.
UpdateMemberRoleByID(ctx context.Context, memberID, teamID, role string) (int64, error)
// DeleteMemberByID deletes a member using the member row ID.
DeleteMemberByID(ctx context.Context, memberID, teamID string) (int64, error)
// ListTeamAuditActions returns distinct audit actions for a team's members.
ListTeamAuditActions(ctx context.Context, teamID string) ([]string, error)
// ── CS5b additions (v0.29.0) ──
// GetFirstTeamIDForUser returns the first team_id the user belongs to, or "" if none.
GetFirstTeamIDForUser(ctx context.Context, userID string) (string, error)
// ── CS6 additions (v0.29.0) ──
// AddMemberReturningID inserts a team member and returns the row ID.
AddMemberReturningID(ctx context.Context, teamID, userID, role string) (string, error)
// HasPrivateProviderRequirement checks if a user belongs to any active team
// with the require_private_providers setting enabled.
HasPrivateProviderRequirement(ctx context.Context, userID string) (bool, error)
// MergeSettings merges a JSON string into the team's settings column.
MergeSettings(ctx context.Context, teamID, settingsJSON string) error
} }
// ========================================= // =========================================
@@ -260,6 +431,169 @@ type ChannelStore interface {
// Admin: archived channel management (v0.23.2) // Admin: archived channel management (v0.23.2)
ListArchived(ctx context.Context, opts ListOptions) ([]ArchivedChannel, int, error) ListArchived(ctx context.Context, opts ListOptions) ([]ArchivedChannel, int, error)
Purge(ctx context.Context, id string) error // hard-delete; channel must be archived Purge(ctx context.Context, id string) error // hard-delete; channel must be archived
// ── Single-field helpers (v0.29.0 — moved from handler raw SQL) ──
// GetAIMode returns the channel's ai_mode ("auto", "off", "mention_only").
GetAIMode(ctx context.Context, channelID string) (string, error)
// GetTypeAndTeamID returns (channel_type, team_id) in one query.
GetTypeAndTeamID(ctx context.Context, channelID string) (string, *string, error)
// GetSystemPrompt returns the channel's system_prompt (may be nil).
GetSystemPrompt(ctx context.Context, channelID string) (*string, error)
// GetDefaultModel returns the channel's model field (may be nil).
GetDefaultModel(ctx context.Context, channelID string) (*string, error)
// TouchUpdatedAt bumps the channel's updated_at to now.
TouchUpdatedAt(ctx context.Context, channelID string) error
// ListUserParticipantIDs returns user participant IDs, excluding excludeUserID.
ListUserParticipantIDs(ctx context.Context, channelID, excludeUserID string) ([]string, error)
// ListPersonaParticipantIDs returns persona participant IDs ordered by created_at.
ListPersonaParticipantIDs(ctx context.Context, channelID string) ([]string, error)
// GetLeaderPersonaID returns the persona group leader for a channel.
// Falls back to the first persona participant if no leader flag is set.
// Returns empty string if no persona participants exist.
GetLeaderPersonaID(ctx context.Context, channelID string) (string, error)
// GetWorkflowInfo returns (workflow_id, current_stage) for workflow-type channels.
// Returns nil workflow_id if the channel is not a workflow channel.
GetWorkflowInfo(ctx context.Context, channelID string) (*string, int, error)
// ── DM + lifecycle helpers (v0.29.0-cs1) ──
// FindExistingDM returns the channel ID of an existing DM between two users.
// Returns "" if no DM exists.
FindExistingDM(ctx context.Context, userID1, userID2 string) (string, error)
// GetUnreadCount returns the number of unread messages for a user in a channel.
GetUnreadCount(ctx context.Context, channelID, userID string) (int, error)
// DeleteByOwner deletes a channel if the user is the owner.
// Returns rows affected (0 = not found or not owner).
DeleteByOwner(ctx context.Context, channelID, userID string) (int64, error)
// MarkRead updates last_read_at and last_read_message_id for a user.
MarkRead(ctx context.Context, channelID, userID string) error
// CountParticipantsByType returns count of participants of a given type.
CountParticipantsByType(ctx context.Context, channelID, pType string) (int, error)
// CountAll returns the total number of channels.
CountAll(ctx context.Context) (int, error)
// ── Workflow instance state (v0.29.0-cs3) ──
// SetWorkflowInstance initializes workflow columns on a channel.
SetWorkflowInstance(ctx context.Context, channelID, workflowID string, version int, stageData json.RawMessage, status string) error
// GetWorkflowStatus returns the full workflow state for a workflow channel.
GetWorkflowStatus(ctx context.Context, channelID string) (*WorkflowChannelStatus, error)
// AdvanceWorkflowStage updates current_stage + stage_data + last_activity_at.
AdvanceWorkflowStage(ctx context.Context, channelID string, nextStage int, stageData json.RawMessage) error
// CompleteWorkflow sets workflow_status='completed' and ai_mode='off'.
CompleteWorkflow(ctx context.Context, channelID string, finalStage int, stageData json.RawMessage) error
// RejectWorkflowToStage resets current_stage (no stage_data change).
RejectWorkflowToStage(ctx context.Context, channelID string, stage int) error
// GetStageData returns the current stage_data JSON from a channel.
GetStageData(ctx context.Context, channelID string) (json.RawMessage, error)
// ── Background job helpers (v0.29.0-cs4) ──
// MarkStaleWorkflows sets workflow_status='stale' on active workflow channels
// with last_activity_at before cutoff. Returns rows affected.
MarkStaleWorkflows(ctx context.Context, cutoff time.Time) (int64, error)
// EnforceWorkflowRetention deletes completed workflow channels whose parent
// workflow has retention.mode="delete" and the retention period has elapsed.
// PG-only (uses JSON operators). Returns rows affected. No-op on SQLite.
EnforceWorkflowRetention(ctx context.Context) (int64, error)
// GetTypeAndAllowAnonymous returns (type, allow_anonymous) for session auth.
GetTypeAndAllowAnonymous(ctx context.Context, channelID string) (string, bool, error)
// ── CS5b additions (v0.29.0) ──
// FindCompactionCandidates returns direct channels eligible for auto-compaction.
// Filters by activity gap, max age, minimum message count, and minimum total chars.
// PG-only (uses settings::text cast). Returns empty slice on SQLite.
FindCompactionCandidates(ctx context.Context, activityBefore, createdAfter time.Time, minMessages, minChars, limit int) ([]models.Channel, error)
// ── CS7a additions (v0.29.0) ──
// GetProviderConfigID returns the channel's provider_config_id (may be nil).
GetProviderConfigID(ctx context.Context, channelID string) (*string, error)
// UserCanAccess checks if a user is the channel owner or a participant.
UserCanAccess(ctx context.Context, channelID, userID string) (bool, error)
// ── CS7b additions (v0.29.0) ──
// ListFiltered returns paginated channels with message counts, applying
// complex filters (types, folder, search, project). Includes channels
// owned by the user or where the user is a participant.
ListFiltered(ctx context.Context, userID string, filter ChannelListFilter) ([]ChannelListItem, int, error)
// GetForUser loads a single channel with message count, verifying the
// user is the owner or a participant. Returns sql.ErrNoRows if not found.
GetForUser(ctx context.Context, channelID, userID string) (*ChannelListItem, error)
// CreateFull atomically creates a channel, adds the owner as a participant,
// optionally adds DM partner participants, and creates a default channel_model.
// The channel model (ch) is populated with ID, CreatedAt, UpdatedAt on return.
// folder, tags, and aiMode are passed separately as they're not in models.Channel.
CreateFull(ctx context.Context, ch *models.Channel, folder string, tags []string, aiMode string,
ownerUserID string, dmPartnerIDs []string, defaultModel, defaultConfigID string) error
// MergeSettings merges a JSON object into the channel's settings column.
MergeSettings(ctx context.Context, channelID string, settingsJSON json.RawMessage) error
}
// ChannelListFilter holds filter options for ChannelStore.ListFiltered.
type ChannelListFilter struct {
ListOptions
Archived bool
Types []string // empty = all
Folder string
FolderID string
Search string // title ILIKE/LIKE
ProjectID string // "none" = NULL, uuid = specific, "" = no filter
}
// ChannelListItem is returned by ListFiltered and GetForUser.
type ChannelListItem struct {
ID string
UserID string
Title string
Type string
AiMode string
Topic *string
Description *string
Model *string
ProviderConfigID *string
SystemPrompt *string
IsArchived bool
IsPinned bool
Folder *string
FolderID *string
ProjectID *string
WorkspaceID *string
Tags []string
Settings json.RawMessage
MessageCount int
UnreadCount int
CreatedAt string
UpdatedAt string
CreatedAtTime time.Time `json:"-"` // SQLite scan target
UpdatedAtTime time.Time `json:"-"` // SQLite scan target
} }
// ArchivedChannel is a view model for the admin archived channels list. // ArchivedChannel is a view model for the admin archived channels list.
@@ -291,6 +625,89 @@ type MessageStore interface {
// Count // Count
CountForChannel(ctx context.Context, channelID string) (int, error) CountForChannel(ctx context.Context, channelID string) (int, error)
// ── Tree operations (v0.29.0 — moved from treepath package) ──
// GetActiveLeaf returns the cursor's active_leaf_id for a user in a channel.
// Falls back to the chronologically latest live message if no cursor exists.
GetActiveLeaf(ctx context.Context, channelID, userID string) (*string, error)
// GetPathToLeaf walks from a leaf up to root via parent_id, returning
// root-first order with sibling counts and sender info resolved.
GetPathToLeaf(ctx context.Context, channelID, leafID string) ([]PathMessage, error)
// GetActivePath is GetActiveLeaf + GetPathToLeaf combined.
GetActivePath(ctx context.Context, channelID, userID string) ([]PathMessage, error)
// GetSiblingsList returns all live siblings of a message with the current index.
GetSiblingsList(ctx context.Context, messageID string) ([]SiblingInfo, int, error)
// GetSiblingCount returns how many live siblings share the same parent.
GetSiblingCount(ctx context.Context, channelID string, parentID *string) (int, error)
// FindLeafFromMessage walks down from a message to find the deepest descendant.
FindLeafFromMessage(ctx context.Context, messageID string) (string, error)
// NextSiblingIndexForParent returns the next sibling_index for children of parentID.
NextSiblingIndexForParent(ctx context.Context, channelID string, parentID *string) (int, error)
// HasPersonaMessages checks if any assistant message in the channel was
// generated by a persona.
HasPersonaMessages(ctx context.Context, channelID string) (bool, error)
// CreateWithCursor inserts a message, updates the cursor, and touches
// channel.updated_at. Handles dialect-specific ID generation internally.
CreateWithCursor(ctx context.Context, m *models.Message, cursorUserID string) error
// ResolveSenderInfo batch-resolves sender_name and sender_avatar for
// user and persona participants. Updates the PathMessage slice in-place.
ResolveSenderInfo(ctx context.Context, path []PathMessage) error
// CountAll returns the total number of messages (for admin stats).
CountAll(ctx context.Context) (int, error)
// ── CS5c additions (v0.29.0) ──
// SearchInChannel performs full-text search within a single channel's messages.
// PG uses to_tsvector/ts_headline; SQLite falls back to LIKE.
// roleFilter is "user", "assistant", or "" for all.
SearchInChannel(ctx context.Context, channelID, query, roleFilter string, limit int) ([]ChannelSearchResult, error)
// ── CS7a additions (v0.29.0) ──
// ListWithSenderInfo returns paginated messages with sender name/avatar
// resolved via JOINs to users and personas tables.
ListWithSenderInfo(ctx context.Context, channelID string, limit, offset int) ([]MessageWithSender, int, error)
// GetParentAndRole loads just the parent_id and role of a message,
// verifying it belongs to the given channel and is not deleted.
GetParentAndRole(ctx context.Context, messageID, channelID string) (parentID *string, role string, err error)
}
// MessageWithSender is returned by MessageStore.ListWithSenderInfo.
type MessageWithSender struct {
ID string
ChannelID string
Role string
Content string
Model *string
TokensUsed *int
ParentID *string
SiblingIndex int
ParticipantType *string
ParticipantID *string
SenderName *string
SenderAvatar *string
CreatedAt string
}
// ChannelSearchResult is returned by MessageStore.SearchInChannel.
type ChannelSearchResult struct {
MessageID string
Role string
Excerpt string // ts_headline or content substring
Rank float64
Timestamp time.Time
} }
// ========================================= // =========================================
@@ -311,6 +728,7 @@ type AuditListOptions struct {
ResourceID string ResourceID string
Since *time.Time Since *time.Time
Until *time.Time Until *time.Time
TeamID string // CS6: scope to team members
} }
// ========================================= // =========================================
@@ -326,6 +744,22 @@ type NoteStore interface {
Search(ctx context.Context, userID, query string, opts ListOptions) ([]models.Note, int, error) Search(ctx context.Context, userID, query string, opts ListOptions) ([]models.Note, int, error)
SearchTitles(ctx context.Context, userID, query string, limit int) ([]models.Note, error) SearchTitles(ctx context.Context, userID, query string, limit int) ([]models.Note, error)
BulkDelete(ctx context.Context, ids []string, userID string) (int, error) BulkDelete(ctx context.Context, ids []string, userID string) (int, error)
// ── CS5c additions (v0.29.0) ──
// SetEmbedding stores a pgvector embedding on a note. PG-only; SQLite no-ops.
SetEmbedding(ctx context.Context, noteID, vecStr string) error
// SearchKeyword performs full-text keyword search with ranking and headlines.
// PG uses ts_rank/ts_headline; SQLite falls back to LIKE.
SearchKeyword(ctx context.Context, userID, query string, limit int) ([]NoteSearchResult, error)
// SearchSemantic performs vector similarity search against note embeddings.
// PG-only; SQLite returns empty results.
SearchSemantic(ctx context.Context, userID, vecStr string, limit int) ([]NoteSearchResult, error)
// ListFolders returns distinct folder paths and note counts for a user.
ListFolders(ctx context.Context, userID string) ([]FolderInfo, error)
} }
type NoteListOptions struct { type NoteListOptions struct {
@@ -335,6 +769,23 @@ type NoteListOptions struct {
TeamID string TeamID string
} }
// NoteSearchResult is returned by NoteStore.SearchKeyword and SearchSemantic.
type NoteSearchResult struct {
ID string
Title string
FolderPath string
Tags []string
Excerpt string // truncated content preview
Headline string // ts_headline markup (empty on SQLite)
Rank float64 // ts_rank or similarity score
}
// FolderInfo is returned by NoteStore.ListFolders.
type FolderInfo struct {
Path string
Count int
}
// ========================================= // =========================================
// NOTE LINK STORE // NOTE LINK STORE
// ========================================= // =========================================
@@ -362,6 +813,24 @@ type GlobalConfigStore interface {
Get(ctx context.Context, key string) (models.JSONMap, error) Get(ctx context.Context, key string) (models.JSONMap, error)
Set(ctx context.Context, key string, value models.JSONMap, updatedBy string) error Set(ctx context.Context, key string, value models.JSONMap, updatedBy string) error
GetAll(ctx context.Context) (map[string]models.JSONMap, error) GetAll(ctx context.Context) (map[string]models.JSONMap, error)
// ── OIDC state (v0.29.0-cs4) ──
// SaveOIDCState stores a state+nonce pair for OIDC callback verification.
SaveOIDCState(ctx context.Context, state, nonce, redirectTo string) error
// ConsumeOIDCState retrieves and deletes an OIDC state (one-time use).
// Returns (nonce, redirectTo, error). Returns sql.ErrNoRows if not found.
ConsumeOIDCState(ctx context.Context, state string) (nonce, redirectTo string, err error)
// CleanupOIDCState removes stale OIDC states older than 10 minutes.
CleanupOIDCState(ctx context.Context) error
// ── CS6 additions (v0.29.0) ──
// GetString returns the raw value column as a string for simple settings
// (e.g. bare JSON booleans like "true"/"false").
GetString(ctx context.Context, key string) (string, error)
} }
// ========================================= // =========================================
@@ -419,6 +888,11 @@ type FileStore interface {
DeleteByChannel(ctx context.Context, channelID string) ([]string, error) // returns storage_keys DeleteByChannel(ctx context.Context, channelID string) ([]string, error) // returns storage_keys
UserUsageBytes(ctx context.Context, userID string) (int64, error) UserUsageBytes(ctx context.Context, userID string) (int64, error)
ListOrphans(ctx context.Context, olderThan time.Duration) ([]models.File, error) ListOrphans(ctx context.Context, olderThan time.Duration) ([]models.File, error)
// ── CS5c additions (v0.29.0) ──
// UpdateStorageKey sets the storage key after initial file creation.
UpdateStorageKey(ctx context.Context, id, key string) error
} }
@@ -652,6 +1126,70 @@ type SessionStore interface {
DeleteExpired(ctx context.Context, olderThan time.Time) (int64, error) DeleteExpired(ctx context.Context, olderThan time.Time) (int64, error)
} }
// =========================================
// PERSONA GROUP STORE (v0.29.0)
// =========================================
// PersonaGroupStore manages persona_groups and their members.
type PersonaGroupStore interface {
// CRUD
List(ctx context.Context, ownerID string) ([]models.PersonaGroup, error)
Get(ctx context.Context, id, ownerID string) (*models.PersonaGroup, error)
Create(ctx context.Context, g *models.PersonaGroup) error
Update(ctx context.Context, id string, fields map[string]interface{}) error
Delete(ctx context.Context, id, ownerID string) (int64, error)
// Ownership check
GetOwnerID(ctx context.Context, id string) (string, error)
// Members
AddMember(ctx context.Context, groupID, personaID string, isLeader bool) error
RemoveMember(ctx context.Context, memberID, groupID string) error
ListMembers(ctx context.Context, groupID string) ([]models.PersonaGroupMember, error)
}
// =========================================
// FOLDER STORE (v0.29.0)
// =========================================
// FolderStore manages user chat folders.
type FolderStore interface {
List(ctx context.Context, userID string) ([]models.Folder, error)
Create(ctx context.Context, f *models.Folder) error
Update(ctx context.Context, folderID, userID string, name string, sortOrder *int) (int64, error)
Delete(ctx context.Context, folderID, userID string) (int64, error)
// UnassignChannels removes folder_id from all channels in this folder.
UnassignChannels(ctx context.Context, folderID, userID string) error
}
// =========================================
// HEALTH STORE (v0.29.0)
// =========================================
// HealthStore manages provider health windows. The full interface is used
// by the health accumulator and status querier; the scheduler only needs Prune.
type HealthStore interface {
// Prune deletes health windows older than the given time. Returns rows affected.
Prune(ctx context.Context, before time.Time) (int64, error)
}
// =========================================
// PRESENCE STORE (v0.29.0)
// =========================================
// PresenceStore manages user online/offline status.
type PresenceStore interface {
// Heartbeat upserts the user's last_seen timestamp to now.
Heartbeat(ctx context.Context, userID string) error
// GetLastSeen returns the user's last heartbeat time, or nil if never seen.
GetLastSeen(ctx context.Context, userID string) (*time.Time, error)
// GetStatuses returns online/offline status for a list of user IDs.
// Users with last_seen after threshold are "online", otherwise "offline".
GetStatuses(ctx context.Context, userIDs []string, threshold time.Time) (map[string]string, error)
}
// ========================================= // =========================================
// SHARED TYPES // SHARED TYPES
// ========================================= // =========================================

View File

@@ -25,6 +25,10 @@ type PackageStore interface {
// SetEnabled toggles a package's enabled state. // SetEnabled toggles a package's enabled state.
SetEnabled(ctx context.Context, id string, enabled bool) error SetEnabled(ctx context.Context, id string, enabled bool) error
// SetStatus transitions a package's lifecycle status.
// Valid statuses: active, pending_review, suspended.
SetStatus(ctx context.Context, id string, status string) error
// Delete removes a non-core package. Core packages cannot be deleted. // Delete removes a non-core package. Core packages cannot be deleted.
Delete(ctx context.Context, id string) error Delete(ctx context.Context, id string) error
@@ -79,6 +83,7 @@ type PackageRegistration struct {
InstalledBy *string `json:"installed_by,omitempty" db:"installed_by"` InstalledBy *string `json:"installed_by,omitempty" db:"installed_by"`
Manifest map[string]any `json:"manifest" db:"manifest"` Manifest map[string]any `json:"manifest" db:"manifest"`
Enabled bool `json:"enabled" db:"enabled"` Enabled bool `json:"enabled" db:"enabled"`
Status string `json:"status" db:"status"`
Source string `json:"source" db:"source"` Source string `json:"source" db:"source"`
InstalledAt string `json:"installed_at" db:"installed_at"` InstalledAt string `json:"installed_at" db:"installed_at"`
UpdatedAt string `json:"updated_at" db:"updated_at"` UpdatedAt string `json:"updated_at" db:"updated_at"`

View File

@@ -30,6 +30,9 @@ func (s *AuditStore) List(ctx context.Context, opts store.AuditListOptions) ([]m
al.ip_address, al.user_agent, al.created_at`, "audit_log al") al.ip_address, al.user_agent, al.created_at`, "audit_log al")
b.Join("LEFT JOIN users u ON al.actor_id = u.id") b.Join("LEFT JOIN users u ON al.actor_id = u.id")
if opts.TeamID != "" {
b.Where("al.actor_id IN (SELECT user_id FROM team_members WHERE team_id = ?)", opts.TeamID)
}
if opts.ActorID != "" { if opts.ActorID != "" {
b.Where("al.actor_id = ?", opts.ActorID) b.Where("al.actor_id = ?", opts.ActorID)
} }

View File

@@ -281,3 +281,58 @@ func scanCatalogEntries(rows *sql.Rows) ([]models.CatalogEntry, error) {
} }
return result, rows.Err() return result, rows.Err()
} }
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *CatalogStore) GetCapabilities(ctx context.Context, modelID, configID string) ([]byte, error) {
var capsJSON []byte
err := DB.QueryRowContext(ctx, `
SELECT capabilities FROM model_catalog
WHERE model_id = $1 AND provider_config_id = $2
`, modelID, configID).Scan(&capsJSON)
if err != nil {
return nil, err
}
return capsJSON, nil
}
func (s *CatalogStore) GetCapabilitiesAny(ctx context.Context, modelID string) ([]byte, error) {
var capsJSON []byte
err := DB.QueryRowContext(ctx, `
SELECT capabilities FROM model_catalog
WHERE model_id = $1 ORDER BY last_synced_at DESC LIMIT 1
`, modelID).Scan(&capsJSON)
if err != nil {
return nil, err
}
return capsJSON, nil
}
// ── CS6 additions (v0.29.0) ─────────────────────────────────────────────
func (s *CatalogStore) ListTeamAvailable(ctx context.Context) ([]store.TeamAvailableModel, error) {
rows, err := DB.QueryContext(ctx, `
SELECT mc.id, mc.model_id, mc.display_name, mc.visibility,
ac.provider, ac.name AS provider_name
FROM model_catalog mc
JOIN provider_configs ac ON mc.provider_config_id = ac.id
WHERE mc.visibility IN ('enabled', 'team')
AND ac.is_active = true AND ac.scope = 'global'
ORDER BY ac.name, mc.model_id
`)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.TeamAvailableModel, 0)
for rows.Next() {
var m store.TeamAvailableModel
if err := rows.Scan(&m.ID, &m.ModelID, &m.DisplayName, &m.Visibility,
&m.Provider, &m.ProviderName); err != nil {
continue
}
results = append(results, m)
}
return results, rows.Err()
}

View File

@@ -0,0 +1,57 @@
package postgres
import (
"context"
"database/sql"
)
// ── Mention resolution (v0.29.0) ────────────────────────────────────────
func (s *CatalogStore) FindEnabledByModelID(ctx context.Context, modelID string) (string, string, error) {
var foundModelID, configID string
err := DB.QueryRowContext(ctx, `
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) = LOWER($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
`, modelID).Scan(&foundModelID, &configID)
if err == sql.ErrNoRows {
return "", "", nil
}
return foundModelID, configID, err
}
func (s *CatalogStore) FindEnabledByModelIDPrefix(ctx context.Context, prefix string) (string, string, int, error) {
var count int
err := DB.QueryRowContext(ctx, `
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 LOWER($1)
AND mc.visibility = 'enabled'
AND pc.is_active = true
`, prefix+"%").Scan(&count)
if err != nil {
return "", "", 0, err
}
if count != 1 {
return "", "", count, nil
}
var modelID, configID string
err = DB.QueryRowContext(ctx, `
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 LOWER($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
`, prefix+"%").Scan(&modelID, &configID)
return modelID, configID, 1, err
}

View File

@@ -6,6 +6,9 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"strings" "strings"
"time"
"github.com/lib/pq"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store" "git.gobha.me/xcaliber/chat-switchboard/store"
@@ -471,3 +474,492 @@ func (s *ChannelStore) Purge(ctx context.Context, id string) error {
_, err = DB.ExecContext(ctx, `DELETE FROM channels WHERE id = $1`, id) _, err = DB.ExecContext(ctx, `DELETE FROM channels WHERE id = $1`, id)
return err return err
} }
// ── CS1 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ChannelStore) FindExistingDM(ctx context.Context, userID1, userID2 string) (string, error) {
var channelID string
err := DB.QueryRowContext(ctx, `
SELECT cp1.channel_id FROM channel_participants cp1
JOIN channel_participants cp2 ON cp1.channel_id = cp2.channel_id
JOIN channels c ON c.id = cp1.channel_id
WHERE c.type = 'dm'
AND cp1.participant_type = 'user' AND cp1.participant_id = $1
AND cp2.participant_type = 'user' AND cp2.participant_id = $2
LIMIT 1
`, userID1, userID2).Scan(&channelID)
if err == sql.ErrNoRows {
return "", nil
}
return channelID, err
}
func (s *ChannelStore) GetUnreadCount(ctx context.Context, channelID, userID string) (int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages m
JOIN channel_participants cp ON cp.channel_id = m.channel_id
WHERE cp.channel_id = $1
AND cp.participant_type = 'user' AND cp.participant_id = $2
AND m.created_at > cp.last_read_at
`, channelID, userID).Scan(&count)
return count, err
}
func (s *ChannelStore) DeleteByOwner(ctx context.Context, channelID, userID string) (int64, error) {
result, err := DB.ExecContext(ctx,
`DELETE FROM channels WHERE id = $1 AND user_id = $2`,
channelID, userID)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *ChannelStore) MarkRead(ctx context.Context, channelID, userID string) error {
// Update last_read_at
_, err := DB.ExecContext(ctx, `
UPDATE channel_participants
SET last_read_at = NOW()
WHERE channel_id = $1 AND participant_type = 'user' AND participant_id = $2
`, channelID, userID)
if err != nil {
return nil // participant may not exist for legacy chats
}
// Best-effort: update last_read_message_id
var latestMsgID *string
_ = DB.QueryRowContext(ctx, `
SELECT id FROM messages WHERE channel_id = $1 ORDER BY created_at DESC LIMIT 1
`, channelID).Scan(&latestMsgID)
if latestMsgID != nil {
_, _ = DB.ExecContext(ctx, `
UPDATE channel_participants
SET last_read_message_id = $1
WHERE channel_id = $2 AND participant_type = 'user' AND participant_id = $3
`, *latestMsgID, channelID, userID)
}
return nil
}
func (s *ChannelStore) CountParticipantsByType(ctx context.Context, channelID, pType string) (int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM channel_participants
WHERE channel_id = $1 AND participant_type = $2
`, channelID, pType).Scan(&count)
return count, err
}
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ChannelStore) CountAll(ctx context.Context) (int, error) {
var count int
err := DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM channels").Scan(&count)
return count, err
}
// ── Workflow instance state (v0.29.0-cs3) ───────────────────────────────
func (s *ChannelStore) SetWorkflowInstance(ctx context.Context, channelID, workflowID string, version int, stageData json.RawMessage, status string) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET workflow_id = $1, workflow_version = $2, current_stage = 0,
stage_data = $3, workflow_status = $4, last_activity_at = $5
WHERE id = $6
`, workflowID, version, stageData, status, time.Now().UTC(), channelID)
return err
}
func (s *ChannelStore) GetWorkflowStatus(ctx context.Context, channelID string) (*store.WorkflowChannelStatus, error) {
var ws store.WorkflowChannelStatus
var stageData []byte
err := DB.QueryRowContext(ctx, `
SELECT workflow_id, workflow_version, current_stage,
COALESCE(stage_data, '{}'), COALESCE(workflow_status, 'active'),
last_activity_at
FROM channels WHERE id = $1 AND type = 'workflow'
`, channelID).Scan(&ws.WorkflowID, &ws.WorkflowVersion,
&ws.CurrentStage, &stageData, &ws.Status, &ws.LastActivityAt)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
ws.StageData = stageData
return &ws, nil
}
func (s *ChannelStore) AdvanceWorkflowStage(ctx context.Context, channelID string, nextStage int, stageData json.RawMessage) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET current_stage = $1, stage_data = $2, last_activity_at = $3
WHERE id = $4
`, nextStage, stageData, time.Now().UTC(), channelID)
return err
}
func (s *ChannelStore) CompleteWorkflow(ctx context.Context, channelID string, finalStage int, stageData json.RawMessage) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET current_stage = $1, workflow_status = 'completed',
stage_data = $2, last_activity_at = $3, ai_mode = 'off'
WHERE id = $4
`, finalStage, stageData, time.Now().UTC(), channelID)
return err
}
func (s *ChannelStore) RejectWorkflowToStage(ctx context.Context, channelID string, stage int) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels SET current_stage = $1, last_activity_at = $2 WHERE id = $3
`, stage, time.Now().UTC(), channelID)
return err
}
func (s *ChannelStore) GetStageData(ctx context.Context, channelID string) (json.RawMessage, error) {
var data json.RawMessage
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(stage_data, '{}') FROM channels WHERE id = $1
`, channelID).Scan(&data)
return data, err
}
// ── Background job helpers (v0.29.0-cs4) ────────────────────────────────
func (s *ChannelStore) MarkStaleWorkflows(ctx context.Context, cutoff time.Time) (int64, error) {
result, err := DB.ExecContext(ctx, `
UPDATE channels
SET workflow_status = 'stale'
WHERE type = 'workflow'
AND workflow_status = 'active'
AND last_activity_at < $1
`, cutoff)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *ChannelStore) EnforceWorkflowRetention(ctx context.Context) (int64, error) {
result, err := DB.ExecContext(ctx, `
DELETE FROM channels
WHERE type = 'workflow'
AND workflow_status IN ('completed', 'archived')
AND workflow_id IS NOT NULL
AND workflow_id IN (
SELECT id FROM workflows
WHERE retention IS NOT NULL
AND retention->>'mode' = 'delete'
AND (retention->>'delete_after_days')::int > 0
AND channels.last_activity_at < now() - ((retention->>'delete_after_days')::int || ' days')::interval
)
`)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *ChannelStore) GetTypeAndAllowAnonymous(ctx context.Context, channelID string) (string, bool, error) {
var chType string
var allowAnon bool
err := DB.QueryRowContext(ctx,
`SELECT type, allow_anonymous FROM channels WHERE id = $1`, channelID).Scan(&chType, &allowAnon)
return chType, allowAnon, err
}
// ── CS5b additions (v0.29.0) ────────────────────────────────────────────
func (s *ChannelStore) FindCompactionCandidates(ctx context.Context, activityBefore, createdAfter time.Time, minMessages, minChars, limit int) ([]models.Channel, error) {
rows, err := DB.QueryContext(ctx, `
SELECT c.id, c.user_id, COALESCE(c.model, ''), COALESCE(c.settings::text, '{}'),
COUNT(m.id) AS msg_count,
COALESCE(SUM(LENGTH(m.content)), 0) AS total_chars
FROM channels c
JOIN messages m ON m.channel_id = c.id AND m.deleted_at IS NULL
WHERE c.type = 'direct'
AND c.is_archived = false
AND c.updated_at < $1
AND c.updated_at > $2
GROUP BY c.id
HAVING COUNT(m.id) >= $3
AND COALESCE(SUM(LENGTH(m.content)), 0) > $4
ORDER BY c.updated_at DESC
LIMIT $5
`, activityBefore, createdAfter, minMessages, minChars, limit)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.Channel
for rows.Next() {
var ch models.Channel
var settingsRaw string
var msgCount, totalChars int
if err := rows.Scan(&ch.ID, &ch.UserID, &ch.Model, &settingsRaw, &msgCount, &totalChars); err != nil {
continue
}
ch.Settings = models.JSONMap{}
_ = json.Unmarshal([]byte(settingsRaw), &ch.Settings)
result = append(result, ch)
}
if result == nil {
result = []models.Channel{}
}
return result, rows.Err()
}
// ── CS7a additions (v0.29.0) ────────────────────────────────────────────
func (s *ChannelStore) GetProviderConfigID(ctx context.Context, channelID string) (*string, error) {
var configID sql.NullString
err := DB.QueryRowContext(ctx,
`SELECT provider_config_id FROM channels WHERE id = $1`, channelID).Scan(&configID)
if err != nil {
return nil, err
}
return NullableStringPtr(configID), nil
}
func (s *ChannelStore) UserCanAccess(ctx context.Context, channelID, userID string) (bool, error) {
var ok bool
err := DB.QueryRowContext(ctx, `
SELECT EXISTS(
SELECT 1 FROM channels WHERE id = $1 AND user_id = $2
UNION ALL
SELECT 1 FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'user' AND participant_id = $2
LIMIT 1
)`, channelID, userID).Scan(&ok)
return ok, err
}
// ── CS7b additions (v0.29.0) ────────────────────────────────────────────
const channelListCols = `c.id, c.user_id, c.title, c.type, c.ai_mode, c.topic,
c.description, c.model, c.provider_config_id,
c.system_prompt, c.is_archived, c.is_pinned, c.folder, c.folder_id, c.project_id, c.workspace_id,
c.tags, c.settings,
COALESCE(mc.cnt, 0) AS message_count,
c.created_at, c.updated_at`
const channelListFrom = `channels c
LEFT JOIN (
SELECT channel_id, COUNT(*) AS cnt FROM messages WHERE deleted_at IS NULL GROUP BY channel_id
) mc ON mc.channel_id = c.id`
func (s *ChannelStore) ListFiltered(ctx context.Context, userID string, f store.ChannelListFilter) ([]store.ChannelListItem, int, error) {
b := NewSelect(channelListCols, channelListFrom)
b.Where("(c.user_id = ? OR c.id IN (SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = ?))", userID, userID)
b.Where("c.is_archived = ?", f.Archived)
if len(f.Types) == 1 {
b.Where("c.type = ?", f.Types[0])
} else if len(f.Types) > 1 {
placeholders := make([]string, len(f.Types))
for i, t := range f.Types {
b.argIdx++
placeholders[i] = fmt.Sprintf("$%d", b.argIdx)
b.args = append(b.args, t)
}
b.WhereRaw("c.type IN (" + strings.Join(placeholders, ",") + ")")
}
if f.Folder != "" {
b.Where("c.folder = ?", f.Folder)
}
if f.FolderID != "" {
b.Where("c.folder_id = ?", f.FolderID)
}
if f.Search != "" {
b.Where("c.title ILIKE ?", "%"+f.Search+"%")
}
if f.ProjectID == "none" {
b.WhereRaw("c.project_id IS NULL")
} else if f.ProjectID != "" {
b.Where("c.project_id = ?", f.ProjectID)
}
// Count
countQ, countArgs := b.CountBuild()
var total int
DB.QueryRowContext(ctx, countQ, countArgs...).Scan(&total)
b.OrderBy("c.is_pinned DESC, c.updated_at", "DESC")
b.Paginate(f.ListOptions)
q, args := b.Build()
rows, err := DB.QueryContext(ctx, q, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
items, err := scanChannelListItems(rows)
if err != nil {
return nil, 0, err
}
// Compute unread counts
for i := range items {
DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages m
JOIN channel_participants cp ON cp.channel_id = m.channel_id
WHERE cp.channel_id = $1
AND cp.participant_type = 'user' AND cp.participant_id = $2
AND m.created_at > cp.last_read_at
AND m.deleted_at IS NULL
`, items[i].ID, userID).Scan(&items[i].UnreadCount)
}
return items, total, nil
}
func (s *ChannelStore) GetForUser(ctx context.Context, channelID, userID string) (*store.ChannelListItem, error) {
row := DB.QueryRowContext(ctx, fmt.Sprintf(`
SELECT %s FROM %s
WHERE c.id = $1 AND (c.user_id = $2 OR c.id IN (
SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = $2
))
`, channelListCols, channelListFrom), channelID, userID)
var item store.ChannelListItem
var tags []byte
var settings []byte
err := row.Scan(
&item.ID, &item.UserID, &item.Title, &item.Type, &item.AiMode, &item.Topic,
&item.Description, &item.Model, &item.ProviderConfigID,
&item.SystemPrompt, &item.IsArchived, &item.IsPinned, &item.Folder, &item.FolderID, &item.ProjectID, &item.WorkspaceID,
&tags, &settings,
&item.MessageCount, &item.CreatedAtTime, &item.UpdatedAtTime,
)
if err != nil {
return nil, err
}
item.Tags = scanTagsBytes(tags)
item.Settings = safeJSONBytes(settings)
item.CreatedAt = item.CreatedAtTime.Format("2006-01-02T15:04:05Z")
item.UpdatedAt = item.UpdatedAtTime.Format("2006-01-02T15:04:05Z")
return &item, nil
}
func (s *ChannelStore) CreateFull(ctx context.Context, ch *models.Channel, folder string, tags []string, aiMode string,
ownerUserID string, dmPartnerIDs []string, defaultModel, defaultConfigID string) error {
tx, err := DB.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if tags == nil {
tags = []string{}
}
// Insert channel
err = tx.QueryRowContext(ctx, `
INSERT INTO channels (user_id, title, type, description, model, system_prompt,
provider_config_id, folder, folder_id, tags, ai_mode)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11)
RETURNING id, created_at, updated_at`,
ch.UserID, ch.Title, ch.Type, ch.Description, ch.Model, ch.SystemPrompt,
models.NullString(ch.ProviderConfigID), folder, models.NullString(ch.FolderID),
pq.Array(tags), aiMode,
).Scan(&ch.ID, &ch.CreatedAt, &ch.UpdatedAt)
if err != nil {
return fmt.Errorf("CreateFull insert channel: %w", err)
}
// Add owner participant
_, err = tx.ExecContext(ctx, `
INSERT INTO channel_participants (channel_id, participant_type, participant_id, role)
VALUES ($1, 'user', $2, 'owner')
ON CONFLICT DO NOTHING`, ch.ID, ownerUserID)
if err != nil {
return fmt.Errorf("CreateFull add owner: %w", err)
}
// Add DM partner participants
for _, pid := range dmPartnerIDs {
if pid == ownerUserID {
continue
}
_, _ = tx.ExecContext(ctx, `
INSERT INTO channel_participants (channel_id, participant_type, participant_id, role)
VALUES ($1, 'user', $2, 'member')
ON CONFLICT DO NOTHING`, ch.ID, pid)
}
// Auto-create channel_model if model specified
if defaultModel != "" {
_, _ = tx.ExecContext(ctx, `
INSERT INTO channel_models (channel_id, model_id, provider_config_id, is_default)
VALUES ($1, $2, $3, true)
ON CONFLICT DO NOTHING`, ch.ID, defaultModel, models.NullString(&defaultConfigID))
}
return tx.Commit()
}
func (s *ChannelStore) MergeSettings(ctx context.Context, channelID string, settingsJSON json.RawMessage) error {
_, err := DB.ExecContext(ctx,
`UPDATE channels SET settings = COALESCE(settings, '{}'::jsonb) || $1::jsonb WHERE id = $2`,
[]byte(settingsJSON), channelID)
return err
}
// ── CS7b helpers ────────────────────────────
func scanChannelListItems(rows *sql.Rows) ([]store.ChannelListItem, error) {
var result []store.ChannelListItem
for rows.Next() {
var item store.ChannelListItem
var tags []byte
var settings []byte
err := rows.Scan(
&item.ID, &item.UserID, &item.Title, &item.Type, &item.AiMode, &item.Topic,
&item.Description, &item.Model, &item.ProviderConfigID,
&item.SystemPrompt, &item.IsArchived, &item.IsPinned, &item.Folder, &item.FolderID, &item.ProjectID, &item.WorkspaceID,
&tags, &settings,
&item.MessageCount, &item.CreatedAtTime, &item.UpdatedAtTime,
)
if err != nil {
return nil, err
}
item.Tags = scanTagsBytes(tags)
item.Settings = safeJSONBytes(settings)
item.CreatedAt = item.CreatedAtTime.Format("2006-01-02T15:04:05Z")
item.UpdatedAt = item.UpdatedAtTime.Format("2006-01-02T15:04:05Z")
result = append(result, item)
}
return result, rows.Err()
}
func scanTagsBytes(b []byte) []string {
if len(b) == 0 {
return []string{}
}
// Try JSON array first (SQLite path), then PG text[] format
var arr []string
if json.Unmarshal(b, &arr) == nil {
if arr == nil {
return []string{}
}
return arr
}
// PG text[] format: {tag1,tag2}
s := strings.TrimPrefix(strings.TrimSuffix(string(b), "}"), "{")
if s == "" {
return []string{}
}
return strings.Split(s, ",")
}
func safeJSONBytes(b []byte) json.RawMessage {
if len(b) == 0 || !json.Valid(b) {
return json.RawMessage("{}")
}
cp := make([]byte, len(b))
copy(cp, b)
return json.RawMessage(cp)
}

View File

@@ -0,0 +1,146 @@
package postgres
import (
"context"
"database/sql"
)
// ── Single-field helpers (v0.29.0) ──────────────────────────────────────
// Moved from handlers/completion.go and handlers/messages.go raw SQL.
func (s *ChannelStore) GetAIMode(ctx context.Context, channelID string) (string, error) {
var aiMode string
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(ai_mode, 'auto') FROM channels WHERE id = $1
`, channelID).Scan(&aiMode)
if err != nil {
return "auto", err
}
return aiMode, nil
}
func (s *ChannelStore) GetTypeAndTeamID(ctx context.Context, channelID string) (string, *string, error) {
var channelType string
var teamID *string
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(type, 'direct'), team_id FROM channels WHERE id = $1
`, channelID).Scan(&channelType, &teamID)
return channelType, teamID, err
}
func (s *ChannelStore) GetSystemPrompt(ctx context.Context, channelID string) (*string, error) {
var prompt *string
err := DB.QueryRowContext(ctx, `
SELECT system_prompt FROM channels WHERE id = $1
`, channelID).Scan(&prompt)
if err == sql.ErrNoRows {
return nil, nil
}
return prompt, err
}
func (s *ChannelStore) GetDefaultModel(ctx context.Context, channelID string) (*string, error) {
var model *string
err := DB.QueryRowContext(ctx, `
SELECT model FROM channels WHERE id = $1
`, channelID).Scan(&model)
if err == sql.ErrNoRows {
return nil, nil
}
return model, err
}
func (s *ChannelStore) TouchUpdatedAt(ctx context.Context, channelID string) error {
_, err := DB.ExecContext(ctx, `UPDATE channels SET updated_at = NOW() WHERE id = $1`, channelID)
return err
}
func (s *ChannelStore) ListUserParticipantIDs(ctx context.Context, channelID, excludeUserID string) ([]string, error) {
rows, err := DB.QueryContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'user' AND participant_id != $2
`, channelID, excludeUserID)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if ids == nil {
ids = []string{}
}
return ids, rows.Err()
}
func (s *ChannelStore) ListPersonaParticipantIDs(ctx context.Context, channelID string) ([]string, error) {
rows, err := DB.QueryContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'persona'
ORDER BY created_at
`, channelID)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if ids == nil {
ids = []string{}
}
return ids, rows.Err()
}
func (s *ChannelStore) GetLeaderPersonaID(ctx context.Context, channelID string) (string, error) {
// Try group leader first
var leaderID string
err := DB.QueryRowContext(ctx, `
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(&leaderID)
if err == nil && leaderID != "" {
return leaderID, nil
}
// Fallback: first persona participant
err = DB.QueryRowContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = $1 AND participant_type = 'persona'
ORDER BY created_at LIMIT 1
`, channelID).Scan(&leaderID)
if err == sql.ErrNoRows {
return "", nil
}
return leaderID, err
}
func (s *ChannelStore) GetWorkflowInfo(ctx context.Context, channelID string) (*string, int, error) {
var workflowID *string
var currentStage int
err := DB.QueryRowContext(ctx, `
SELECT workflow_id, COALESCE(current_stage, 0)
FROM channels WHERE id = $1 AND type = 'workflow'
`, channelID).Scan(&workflowID, &currentStage)
if err == sql.ErrNoRows {
return nil, 0, nil
}
return workflowID, currentStage, err
}

View File

@@ -0,0 +1,169 @@
package postgres
import (
"context"
"database/sql"
"time"
"git.gobha.me/xcaliber/chat-switchboard/models"
)
type ExtensionPermissionStore struct {
db *sql.DB
}
func NewExtensionPermissionStore(db *sql.DB) *ExtensionPermissionStore {
return &ExtensionPermissionStore{db: db}
}
func (s *ExtensionPermissionStore) DeclareForPackage(ctx context.Context, packageID string, permissions []string) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
// Remove permissions no longer in manifest
if len(permissions) == 0 {
_, err = tx.ExecContext(ctx, `DELETE FROM extension_permissions WHERE package_id = $1`, packageID)
if err != nil {
return err
}
return tx.Commit()
}
// Build set of current declared
declared := make(map[string]bool, len(permissions))
for _, p := range permissions {
declared[p] = true
}
// Get existing
rows, err := tx.QueryContext(ctx,
`SELECT permission FROM extension_permissions WHERE package_id = $1`, packageID)
if err != nil {
return err
}
existing := make(map[string]bool)
for rows.Next() {
var perm string
if err := rows.Scan(&perm); err != nil {
rows.Close()
return err
}
existing[perm] = true
}
rows.Close()
// Delete removed
for perm := range existing {
if !declared[perm] {
_, err = tx.ExecContext(ctx,
`DELETE FROM extension_permissions WHERE package_id = $1 AND permission = $2`,
packageID, perm)
if err != nil {
return err
}
}
}
// Upsert new (preserving existing grants)
for _, perm := range permissions {
if !existing[perm] {
_, err = tx.ExecContext(ctx,
`INSERT INTO extension_permissions (id, package_id, permission)
VALUES (gen_random_uuid(), $1, $2)
ON CONFLICT (package_id, permission) DO NOTHING`,
packageID, perm)
if err != nil {
return err
}
}
}
return tx.Commit()
}
func (s *ExtensionPermissionStore) ListForPackage(ctx context.Context, packageID string) ([]models.ExtensionPermission, error) {
rows, err := s.db.QueryContext(ctx,
`SELECT id, package_id, permission, granted, granted_by, granted_at, created_at
FROM extension_permissions
WHERE package_id = $1
ORDER BY permission`, packageID)
if err != nil {
return nil, err
}
defer rows.Close()
var perms []models.ExtensionPermission
for rows.Next() {
var p models.ExtensionPermission
if err := rows.Scan(&p.ID, &p.PackageID, &p.Permission, &p.Granted, &p.GrantedBy, &p.GrantedAt, &p.CreatedAt); err != nil {
return nil, err
}
perms = append(perms, p)
}
if perms == nil {
perms = []models.ExtensionPermission{}
}
return perms, nil
}
func (s *ExtensionPermissionStore) GrantedForPackage(ctx context.Context, packageID string) ([]string, error) {
rows, err := s.db.QueryContext(ctx,
`SELECT permission FROM extension_permissions
WHERE package_id = $1 AND granted = true
ORDER BY permission`, packageID)
if err != nil {
return nil, err
}
defer rows.Close()
var perms []string
for rows.Next() {
var p string
if err := rows.Scan(&p); err != nil {
return nil, err
}
perms = append(perms, p)
}
if perms == nil {
perms = []string{}
}
return perms, nil
}
func (s *ExtensionPermissionStore) Grant(ctx context.Context, packageID, permission, grantedBy string) error {
now := time.Now()
_, err := s.db.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = true, granted_by = $1, granted_at = $2
WHERE package_id = $3 AND permission = $4`,
grantedBy, now, packageID, permission)
return err
}
func (s *ExtensionPermissionStore) Revoke(ctx context.Context, packageID, permission string) error {
_, err := s.db.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = false, granted_by = NULL, granted_at = NULL
WHERE package_id = $1 AND permission = $2`,
packageID, permission)
return err
}
func (s *ExtensionPermissionStore) GrantAll(ctx context.Context, packageID, grantedBy string) error {
now := time.Now()
_, err := s.db.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = true, granted_by = $1, granted_at = $2
WHERE package_id = $3 AND granted = false`,
grantedBy, now, packageID)
return err
}
func (s *ExtensionPermissionStore) DeleteForPackage(ctx context.Context, packageID string) error {
_, err := s.db.ExecContext(ctx,
`DELETE FROM extension_permissions WHERE package_id = $1`, packageID)
return err
}

View File

@@ -249,3 +249,12 @@ func (s *FileStore) ListOrphans(ctx context.Context, olderThan time.Duration) ([
} }
return out, rows.Err() return out, rows.Err()
} }
// ── CS5c additions (v0.29.0) ──────────────────────────────────────────
func (s *FileStore) UpdateStorageKey(ctx context.Context, id, key string) error {
_, err := DB.ExecContext(ctx,
`UPDATE files SET storage_key = $1, updated_at = NOW() WHERE id = $2`,
key, id)
return err
}

View File

@@ -0,0 +1,75 @@
package postgres
import (
"context"
"git.gobha.me/xcaliber/chat-switchboard/models"
)
type FolderStore struct{}
func NewFolderStore() *FolderStore { return &FolderStore{} }
func (s *FolderStore) List(ctx context.Context, userID string) ([]models.Folder, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, name, parent_id, sort_order, created_at, updated_at
FROM folders WHERE user_id = $1
ORDER BY sort_order, name
`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.Folder
for rows.Next() {
var f models.Folder
if err := rows.Scan(&f.ID, &f.Name, &f.ParentID, &f.SortOrder,
&f.CreatedAt, &f.UpdatedAt); err != nil {
continue
}
f.UserID = userID
result = append(result, f)
}
if result == nil {
result = []models.Folder{}
}
return result, rows.Err()
}
func (s *FolderStore) Create(ctx context.Context, f *models.Folder) error {
return DB.QueryRowContext(ctx, `
INSERT INTO folders (user_id, name, sort_order)
VALUES ($1, $2, $3)
RETURNING id, name, parent_id, sort_order, created_at, updated_at
`, f.UserID, f.Name, f.SortOrder).Scan(
&f.ID, &f.Name, &f.ParentID, &f.SortOrder, &f.CreatedAt, &f.UpdatedAt)
}
func (s *FolderStore) Update(ctx context.Context, folderID, userID string, name string, sortOrder *int) (int64, error) {
res, err := DB.ExecContext(ctx, `
UPDATE folders
SET name = COALESCE(NULLIF($3, ''), name),
sort_order = COALESCE($4, sort_order)
WHERE id = $1 AND user_id = $2
`, folderID, userID, name, sortOrder)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *FolderStore) Delete(ctx context.Context, folderID, userID string) (int64, error) {
res, err := DB.ExecContext(ctx,
`DELETE FROM folders WHERE id = $1 AND user_id = $2`, folderID, userID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *FolderStore) UnassignChannels(ctx context.Context, folderID, userID string) error {
_, err := DB.ExecContext(ctx,
`UPDATE channels SET folder_id = NULL WHERE folder_id = $1 AND user_id = $2`, folderID, userID)
return err
}

View File

@@ -57,3 +57,40 @@ func (s *GlobalConfigStore) GetAll(ctx context.Context) (map[string]models.JSONM
} }
return result, rows.Err() return result, rows.Err()
} }
// ── OIDC state (v0.29.0-cs4) ────────────────────────────────────────────
func (s *GlobalConfigStore) SaveOIDCState(ctx context.Context, state, nonce, redirectTo string) error {
_, err := DB.ExecContext(ctx, `
INSERT INTO oidc_auth_state (state, nonce, redirect_to) VALUES ($1, $2, $3)
`, state, nonce, redirectTo)
return err
}
func (s *GlobalConfigStore) ConsumeOIDCState(ctx context.Context, state string) (string, string, error) {
var nonce, redirectTo string
err := DB.QueryRowContext(ctx, `
SELECT nonce, COALESCE(redirect_to, '') FROM oidc_auth_state WHERE state = $1
`, state).Scan(&nonce, &redirectTo)
if err != nil {
return "", "", err
}
// Delete (one-time use)
_, _ = DB.ExecContext(ctx, `DELETE FROM oidc_auth_state WHERE state = $1`, state)
return nonce, redirectTo, nil
}
func (s *GlobalConfigStore) CleanupOIDCState(ctx context.Context) error {
_, err := DB.ExecContext(ctx,
`DELETE FROM oidc_auth_state WHERE created_at < NOW() - INTERVAL '10 minutes'`)
return err
}
// ── CS6 additions (v0.29.0) ─────────────────────────────────────────────
func (s *GlobalConfigStore) GetString(ctx context.Context, key string) (string, error) {
var val string
err := DB.QueryRowContext(ctx,
"SELECT value FROM global_settings WHERE key = $1", key).Scan(&val)
return val, err
}

View File

@@ -270,3 +270,34 @@ func scanMemories(rows *sql.Rows) ([]models.Memory, error) {
// ensure compile-time interface satisfaction // ensure compile-time interface satisfaction
var _ store.MemoryStore = (*MemoryStore)(nil) var _ store.MemoryStore = (*MemoryStore)(nil)
// ── CS5b additions (v0.29.0) ────────────────────────────────────────────
func (s *MemoryStore) SetEmbedding(ctx context.Context, id, embedding string) error {
_, err := DB.ExecContext(ctx,
`UPDATE memories SET embedding = $1::vector WHERE id = $2`, embedding, id)
return err
}
func (s *MemoryStore) GetLastExtractionMessageID(ctx context.Context, channelID, userID string) (string, error) {
var lastID string
err := DB.QueryRowContext(ctx,
`SELECT last_message_id FROM memory_extraction_log WHERE channel_id = $1 AND user_id = $2`,
channelID, userID).Scan(&lastID)
if err != nil {
return "", nil // no entry yet
}
return lastID, nil
}
func (s *MemoryStore) UpsertExtractionLog(ctx context.Context, channelID, userID, lastMessageID string, count int) error {
_, err := DB.ExecContext(ctx, `
INSERT INTO memory_extraction_log (channel_id, user_id, last_message_id, memory_count)
VALUES ($1, $2, $3, $4)
ON CONFLICT(channel_id, user_id) DO UPDATE SET
last_message_id = EXCLUDED.last_message_id,
extracted_at = now(),
memory_count = memory_extraction_log.memory_count + EXCLUDED.memory_count
`, channelID, userID, lastMessageID, count)
return err
}

View File

@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"fmt"
"time" "time"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
@@ -136,7 +137,9 @@ func (s *MessageStore) GetPathToRoot(ctx context.Context, messageID string) ([]m
m.parent_id, m.sibling_index, m.participant_type, m.participant_id, m.deleted_at, m.created_at m.parent_id, m.sibling_index, m.participant_type, m.participant_id, m.deleted_at, m.created_at
FROM messages m JOIN path p ON m.id = p.parent_id FROM messages m JOIN path p ON m.id = p.parent_id
) )
SELECT * FROM path ORDER BY created_at ASC`, messageID) SELECT id, channel_id, role, content, model, tokens_used, tool_calls, metadata,
parent_id, sibling_index, participant_type, participant_id, deleted_at, created_at
FROM path ORDER BY created_at ASC`, messageID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -189,3 +192,115 @@ func scanMessages(rows *sql.Rows) ([]models.Message, error) {
} }
return result, rows.Err() return result, rows.Err()
} }
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *MessageStore) CountAll(ctx context.Context) (int, error) {
var count int
err := DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM messages").Scan(&count)
return count, err
}
// ── CS5c additions (v0.29.0) ─────────────────────────────────────────────
func (s *MessageStore) SearchInChannel(ctx context.Context, channelID, query, roleFilter string, limit int) ([]store.ChannelSearchResult, error) {
roleClause := ""
queryArgs := []interface{}{channelID, query, limit}
if roleFilter == "user" || roleFilter == "assistant" {
roleClause = "AND m.role = $4"
queryArgs = append(queryArgs, roleFilter)
}
rows, err := DB.QueryContext(ctx, fmt.Sprintf(`
SELECT m.id, m.role,
ts_headline('english', m.content, plainto_tsquery('english', $2),
'MaxWords=60, MinWords=20, StartSel=**, StopSel=**') AS headline,
ts_rank(to_tsvector('english', m.content), plainto_tsquery('english', $2)) AS rank,
m.created_at
FROM messages m
WHERE m.channel_id = $1
AND m.deleted_at IS NULL
AND m.role IN ('user', 'assistant')
AND to_tsvector('english', m.content) @@ plainto_tsquery('english', $2)
%s
ORDER BY rank DESC, m.created_at DESC
LIMIT $3
`, roleClause), queryArgs...)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.ChannelSearchResult, 0)
for rows.Next() {
var r store.ChannelSearchResult
if err := rows.Scan(&r.MessageID, &r.Role, &r.Excerpt, &r.Rank, &r.Timestamp); err != nil {
continue
}
results = append(results, r)
}
return results, rows.Err()
}
// ── CS7a additions (v0.29.0) ────────────────────────────────────────────
func (s *MessageStore) ListWithSenderInfo(ctx context.Context, channelID string, limit, offset int) ([]store.MessageWithSender, int, error) {
var total int
err := DB.QueryRowContext(ctx,
`SELECT COUNT(*) FROM messages WHERE channel_id = $1 AND deleted_at IS NULL`,
channelID).Scan(&total)
if err != nil {
return nil, 0, err
}
rows, err := DB.QueryContext(ctx, `
SELECT m.id, m.channel_id, m.role, m.content, m.model, m.tokens_used, m.parent_id,
m.sibling_index, m.participant_type, m.participant_id,
CASE WHEN m.participant_type = 'user' THEN COALESCE(u.display_name, u.username)
WHEN m.participant_type = 'persona' THEN p.name
ELSE NULL END AS sender_name,
CASE WHEN m.participant_type = 'user' THEN u.avatar_url
WHEN m.participant_type = 'persona' THEN p.avatar
ELSE NULL END AS sender_avatar,
m.created_at
FROM messages m
LEFT JOIN users u ON m.participant_type = 'user' AND m.participant_id = u.id::text
LEFT JOIN personas p ON m.participant_type = 'persona' AND m.participant_id = p.id::text
WHERE m.channel_id = $1 AND m.deleted_at IS NULL
ORDER BY m.created_at ASC
LIMIT $2 OFFSET $3
`, channelID, limit, offset)
if err != nil {
return nil, 0, err
}
defer rows.Close()
results := make([]store.MessageWithSender, 0)
for rows.Next() {
var m store.MessageWithSender
if err := rows.Scan(
&m.ID, &m.ChannelID, &m.Role, &m.Content,
&m.Model, &m.TokensUsed, &m.ParentID,
&m.SiblingIndex, &m.ParticipantType, &m.ParticipantID,
&m.SenderName, &m.SenderAvatar,
&m.CreatedAt,
); err != nil {
return nil, total, err
}
results = append(results, m)
}
return results, total, rows.Err()
}
func (s *MessageStore) GetParentAndRole(ctx context.Context, messageID, channelID string) (*string, string, error) {
var parentID sql.NullString
var role string
err := DB.QueryRowContext(ctx, `
SELECT parent_id, role FROM messages
WHERE id = $1 AND channel_id = $2 AND deleted_at IS NULL
`, messageID, channelID).Scan(&parentID, &role)
if err != nil {
return nil, "", err
}
return NullableStringPtr(parentID), role, nil
}

View File

@@ -0,0 +1,383 @@
package postgres
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ── Tree Operations (v0.29.0) ───────────────────────────────────────────
// Moved from treepath package. All message tree traversal goes through
// the store interface now.
func (s *MessageStore) GetActiveLeaf(ctx context.Context, channelID, userID string) (*string, error) {
var leafID *string
// Try cursor first
err := DB.QueryRowContext(ctx, `
SELECT active_leaf_id FROM channel_cursors
WHERE channel_id = $1 AND user_id = $2
`, channelID, userID).Scan(&leafID)
if err == nil && leafID != nil {
// Verify the leaf still exists and isn't deleted
var exists bool
DB.QueryRowContext(ctx, `
SELECT EXISTS(SELECT 1 FROM messages WHERE id = $1 AND deleted_at IS NULL)
`, *leafID).Scan(&exists)
if exists {
return leafID, nil
}
}
// Fallback: latest live message in channel
var fallbackID string
err = DB.QueryRowContext(ctx, `
SELECT id FROM messages
WHERE channel_id = $1 AND deleted_at IS NULL
ORDER BY created_at DESC LIMIT 1
`, channelID).Scan(&fallbackID)
if err == sql.ErrNoRows {
return nil, nil // empty channel
}
if err != nil {
return nil, err
}
return &fallbackID, nil
}
func (s *MessageStore) GetPathToLeaf(ctx context.Context, channelID, leafID string) ([]store.PathMessage, error) {
rows, err := DB.QueryContext(ctx, `
WITH RECURSIVE path AS (
SELECT id, parent_id, role, content, model, tokens_used, tool_calls, metadata,
participant_type, participant_id, sibling_index, created_at,
0 AS depth
FROM messages
WHERE id = $1 AND channel_id = $2 AND deleted_at IS NULL
UNION ALL
SELECT m.id, m.parent_id, m.role, m.content, m.model, m.tokens_used, m.tool_calls, m.metadata,
m.participant_type, m.participant_id, m.sibling_index, m.created_at,
p.depth + 1
FROM messages m
JOIN path p ON m.id = p.parent_id
WHERE m.deleted_at IS NULL
)
SELECT id, parent_id, role, content, model, tokens_used, tool_calls, metadata,
participant_type, participant_id, sibling_index, created_at
FROM path
ORDER BY depth DESC
`, leafID, channelID)
if err != nil {
return nil, fmt.Errorf("GetPathToLeaf: %w", err)
}
defer rows.Close()
var path []store.PathMessage
for rows.Next() {
var m store.PathMessage
var participantType, participantID sql.NullString
var toolCallsJSON, metadataJSON []byte
if err := rows.Scan(
&m.ID, &m.ParentID, &m.Role, &m.Content, &m.Model, &m.TokensUsed,
&toolCallsJSON, &metadataJSON,
&participantType, &participantID, &m.SiblingIndex, &m.CreatedAt,
); err != nil {
return nil, fmt.Errorf("GetPathToLeaf scan: %w", err)
}
if len(toolCallsJSON) > 0 && string(toolCallsJSON) != "null" {
raw := json.RawMessage(toolCallsJSON)
m.ToolCalls = &raw
}
if len(metadataJSON) > 0 && string(metadataJSON) != "null" && string(metadataJSON) != "{}" {
raw := json.RawMessage(metadataJSON)
m.Metadata = &raw
}
if participantType.Valid {
m.ParticipantType = participantType.String
}
if participantID.Valid {
m.ParticipantID = participantID.String
}
path = append(path, m)
}
if err := rows.Err(); err != nil {
return nil, err
}
// Enrich with sibling counts
for i := range path {
count, _ := s.GetSiblingCount(ctx, channelID, path[i].ParentID)
path[i].SiblingCount = count
}
// Resolve sender info
_ = s.ResolveSenderInfo(ctx, path)
return path, nil
}
func (s *MessageStore) GetActivePath(ctx context.Context, channelID, userID string) ([]store.PathMessage, error) {
leafID, err := s.GetActiveLeaf(ctx, channelID, userID)
if err != nil {
return nil, err
}
if leafID == nil {
return []store.PathMessage{}, nil
}
return s.GetPathToLeaf(ctx, channelID, *leafID)
}
func (s *MessageStore) GetSiblingsList(ctx context.Context, messageID string) ([]store.SiblingInfo, int, error) {
// Get parent_id and channel_id of the target message
var parentID *string
var channelID string
err := DB.QueryRowContext(ctx, `
SELECT parent_id, channel_id FROM messages
WHERE id = $1 AND deleted_at IS NULL
`, messageID).Scan(&parentID, &channelID)
if err != nil {
return nil, 0, fmt.Errorf("message not found: %w", err)
}
var rows *sql.Rows
if parentID == nil {
rows, err = DB.QueryContext(ctx, `
SELECT id, role, model, sibling_index, SUBSTR(content, 1, 80), created_at
FROM messages
WHERE channel_id = $1 AND parent_id IS NULL AND deleted_at IS NULL
ORDER BY sibling_index, created_at
`, channelID)
} else {
rows, err = DB.QueryContext(ctx, `
SELECT id, role, model, sibling_index, SUBSTR(content, 1, 80), created_at
FROM messages
WHERE parent_id = $1 AND deleted_at IS NULL
ORDER BY sibling_index, created_at
`, *parentID)
}
if err != nil {
return nil, 0, err
}
defer rows.Close()
var siblings []store.SiblingInfo
currentIdx := 0
for i := 0; rows.Next(); i++ {
var si store.SiblingInfo
if err := rows.Scan(&si.ID, &si.Role, &si.Model, &si.SiblingIndex, &si.Preview, &si.CreatedAt); err != nil {
return nil, 0, err
}
if si.ID == messageID {
currentIdx = i
}
siblings = append(siblings, si)
}
return siblings, currentIdx, rows.Err()
}
func (s *MessageStore) GetSiblingCount(ctx context.Context, channelID string, parentID *string) (int, error) {
var count int
var err error
if parentID == nil {
err = DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages
WHERE channel_id = $1 AND parent_id IS NULL AND deleted_at IS NULL
`, channelID).Scan(&count)
} else {
err = DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages
WHERE parent_id = $1 AND deleted_at IS NULL
`, *parentID).Scan(&count)
}
if err != nil || count == 0 {
return 1, nil // minimum 1 (the message itself)
}
return count, nil
}
func (s *MessageStore) FindLeafFromMessage(ctx context.Context, messageID string) (string, error) {
var leafID string
err := DB.QueryRowContext(ctx, `
WITH RECURSIVE descendants AS (
SELECT id, 0 AS depth
FROM messages
WHERE id = $1 AND deleted_at IS NULL
UNION ALL
SELECT child.id, d.depth + 1
FROM messages child
JOIN descendants d ON child.parent_id = d.id
WHERE child.deleted_at IS NULL
AND child.sibling_index = (
SELECT MIN(sibling_index) FROM messages
WHERE parent_id = d.id AND deleted_at IS NULL
)
)
SELECT id FROM descendants
ORDER BY depth DESC
LIMIT 1
`, messageID).Scan(&leafID)
if err != nil {
return messageID, nil // fallback to message itself
}
return leafID, nil
}
func (s *MessageStore) NextSiblingIndexForParent(ctx context.Context, channelID string, parentID *string) (int, error) {
var maxIdx sql.NullInt64
var err error
if parentID == nil {
err = DB.QueryRowContext(ctx, `
SELECT MAX(sibling_index) FROM messages
WHERE channel_id = $1 AND parent_id IS NULL AND deleted_at IS NULL
`, channelID).Scan(&maxIdx)
} else {
err = DB.QueryRowContext(ctx, `
SELECT MAX(sibling_index) FROM messages
WHERE parent_id = $1 AND deleted_at IS NULL
`, *parentID).Scan(&maxIdx)
}
if err != nil || !maxIdx.Valid {
return 0, nil
}
return int(maxIdx.Int64) + 1, nil
}
func (s *MessageStore) HasPersonaMessages(ctx context.Context, channelID string) (bool, error) {
var id string
err := DB.QueryRowContext(ctx, `
SELECT id FROM messages
WHERE channel_id = $1 AND role = 'assistant' AND participant_type = 'persona'
LIMIT 1
`, channelID).Scan(&id)
return err == nil && id != "", nil
}
func (s *MessageStore) CreateWithCursor(ctx context.Context, m *models.Message, cursorUserID string) error {
// Insert message — PG generates ID via gen_random_uuid()
err := DB.QueryRowContext(ctx, `
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, created_at`,
m.ChannelID, m.Role, m.Content, safeModel(m.Model), m.TokensUsed,
ToJSON(m.ToolCalls), models.NullString(m.ParentID),
m.ParticipantType, m.ParticipantID,
safeString(m.ProviderConfigID), m.SiblingIndex,
).Scan(&m.ID, &m.CreatedAt)
if err != nil {
return fmt.Errorf("CreateWithCursor insert: %w", err)
}
// Update cursor
if cursorUserID != "" {
_, _ = DB.ExecContext(ctx, `
INSERT INTO channel_cursors (channel_id, user_id, active_leaf_id)
VALUES ($1, $2, $3)
ON CONFLICT (channel_id, user_id) DO UPDATE SET active_leaf_id = $3, updated_at = NOW()
`, m.ChannelID, cursorUserID, m.ID)
}
// Touch channel
_, _ = DB.ExecContext(ctx, `UPDATE channels SET updated_at = NOW() WHERE id = $1`, m.ChannelID)
return nil
}
func (s *MessageStore) ResolveSenderInfo(ctx context.Context, path []store.PathMessage) error {
// Collect unique participant IDs by type
userIDs := map[string]bool{}
personaIDs := map[string]bool{}
for _, m := range path {
if m.ParticipantID == "" {
continue
}
switch m.ParticipantType {
case "user":
userIDs[m.ParticipantID] = true
case "persona":
personaIDs[m.ParticipantID] = true
}
}
// Resolve users
userNames := map[string]string{}
userAvatars := map[string]string{}
for uid := range userIDs {
var name, avatar sql.NullString
_ = DB.QueryRowContext(ctx, `
SELECT COALESCE(display_name, username), avatar_url FROM users WHERE id = $1
`, uid).Scan(&name, &avatar)
if name.Valid {
userNames[uid] = name.String
}
if avatar.Valid {
userAvatars[uid] = avatar.String
}
}
// Resolve personas
personaNames := map[string]string{}
personaAvatars := map[string]string{}
for pid := range personaIDs {
var name, avatar sql.NullString
_ = DB.QueryRowContext(ctx, `
SELECT name, avatar FROM personas WHERE id = $1
`, pid).Scan(&name, &avatar)
if name.Valid {
personaNames[pid] = name.String
}
if avatar.Valid {
personaAvatars[pid] = avatar.String
}
}
// Apply to path
for i := range path {
pid := path[i].ParticipantID
switch path[i].ParticipantType {
case "user":
if n, ok := userNames[pid]; ok {
path[i].SenderName = &n
}
if a, ok := userAvatars[pid]; ok && a != "" {
path[i].SenderAvatar = &a
}
case "persona":
if n, ok := personaNames[pid]; ok {
path[i].SenderName = &n
}
if a, ok := personaAvatars[pid]; ok && a != "" {
path[i].SenderAvatar = &a
}
}
}
return nil
}
// ── helpers ─────────────────────────────────
func safeModel(m string) interface{} {
if m == "" {
return ""
}
return m
}
func safeString(s *string) interface{} {
if s == nil || *s == "" {
return nil
}
return *s
}

View File

@@ -219,3 +219,100 @@ func (s *NoteStore) SearchTitles(ctx context.Context, userID, query string, limi
} }
return results, rows.Err() return results, rows.Err()
} }
// ── CS5c additions (v0.29.0) ──────────────────────────────────────────
func (s *NoteStore) SetEmbedding(ctx context.Context, noteID, vecStr string) error {
_, err := DB.ExecContext(ctx,
`UPDATE notes SET embedding = $1::vector WHERE id = $2`,
vecStr, noteID)
return err
}
func (s *NoteStore) SearchKeyword(ctx context.Context, userID, query string, limit int) ([]store.NoteSearchResult, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, title, folder_path, tags, LEFT(content, 500),
ts_rank(search_vector, plainto_tsquery('english', $2)) AS rank,
ts_headline('english', content, plainto_tsquery('english', $2),
'MaxWords=60, MinWords=20, StartSel=**, StopSel=**') AS headline
FROM notes
WHERE user_id = $1
AND search_vector @@ plainto_tsquery('english', $2)
ORDER BY rank DESC
LIMIT $3
`, userID, query, limit)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.NoteSearchResult, 0)
for rows.Next() {
var r store.NoteSearchResult
var dbTags pq.StringArray
if err := rows.Scan(&r.ID, &r.Title, &r.FolderPath, &dbTags, &r.Excerpt, &r.Rank, &r.Headline); err != nil {
continue
}
r.Tags = []string(dbTags)
if r.Tags == nil {
r.Tags = []string{}
}
results = append(results, r)
}
return results, rows.Err()
}
func (s *NoteStore) SearchSemantic(ctx context.Context, userID, vecStr string, limit int) ([]store.NoteSearchResult, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, title, folder_path, tags, LEFT(content, 500),
1 - (embedding <=> $2::vector) AS similarity
FROM notes
WHERE user_id = $1
AND embedding IS NOT NULL
AND 1 - (embedding <=> $2::vector) > 0.3
ORDER BY embedding <=> $2::vector
LIMIT $3
`, userID, vecStr, limit)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.NoteSearchResult, 0)
for rows.Next() {
var r store.NoteSearchResult
var dbTags pq.StringArray
if err := rows.Scan(&r.ID, &r.Title, &r.FolderPath, &dbTags, &r.Excerpt, &r.Rank); err != nil {
continue
}
r.Tags = []string(dbTags)
if r.Tags == nil {
r.Tags = []string{}
}
results = append(results, r)
}
return results, rows.Err()
}
func (s *NoteStore) ListFolders(ctx context.Context, userID string) ([]store.FolderInfo, error) {
rows, err := DB.QueryContext(ctx, `
SELECT DISTINCT folder_path, COUNT(*) AS count
FROM notes WHERE user_id = $1
GROUP BY folder_path
ORDER BY folder_path
`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.FolderInfo, 0)
for rows.Next() {
var f store.FolderInfo
if err := rows.Scan(&f.Path, &f.Count); err != nil {
continue
}
results = append(results, f)
}
return results, rows.Err()
}

View File

@@ -51,6 +51,20 @@ func (s *PackageStore) SetEnabled(ctx context.Context, id string, enabled bool)
return nil return nil
} }
func (s *PackageStore) SetStatus(ctx context.Context, id string, status string) error {
result, err := DB.ExecContext(ctx,
`UPDATE packages SET status = $2, updated_at = NOW() WHERE id = $1`,
id, status)
if err != nil {
return err
}
affected, _ := result.RowsAffected()
if affected == 0 {
return sql.ErrNoRows
}
return nil
}
func (s *PackageStore) Delete(ctx context.Context, id string) error { func (s *PackageStore) Delete(ctx context.Context, id string) error {
result, err := DB.ExecContext(ctx, result, err := DB.ExecContext(ctx,
`DELETE FROM packages WHERE id = $1 AND source != 'core'`, id) `DELETE FROM packages WHERE id = $1 AND source != 'core'`, id)
@@ -87,15 +101,18 @@ func (s *PackageStore) ListEnabled(ctx context.Context) ([]string, error) {
func (s *PackageStore) Create(ctx context.Context, pkg *store.PackageRegistration) error { func (s *PackageStore) Create(ctx context.Context, pkg *store.PackageRegistration) error {
manifestJSON := ToJSON(pkg.Manifest) manifestJSON := ToJSON(pkg.Manifest)
if pkg.Status == "" {
pkg.Status = "active"
}
return DB.QueryRowContext(ctx, ` return DB.QueryRowContext(ctx, `
INSERT INTO packages (id, title, type, version, description, author, tier, INSERT INTO packages (id, title, type, version, description, author, tier,
is_system, scope, team_id, installed_by, manifest, enabled, source) is_system, scope, team_id, installed_by, manifest, enabled, status, source)
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15)
RETURNING installed_at, updated_at`, RETURNING installed_at, updated_at`,
pkg.ID, pkg.Title, pkg.Type, pkg.Version, pkg.Description, pkg.Author, pkg.ID, pkg.Title, pkg.Type, pkg.Version, pkg.Description, pkg.Author,
pkg.Tier, pkg.IsSystem, pkg.Scope, pkg.Tier, pkg.IsSystem, pkg.Scope,
nullStrPtr(pkg.TeamID), nullStrPtr(pkg.InstalledBy), nullStrPtr(pkg.TeamID), nullStrPtr(pkg.InstalledBy),
manifestJSON, pkg.Enabled, pkg.Source, manifestJSON, pkg.Enabled, pkg.Status, pkg.Source,
).Scan(&pkg.InstalledAt, &pkg.UpdatedAt) ).Scan(&pkg.InstalledAt, &pkg.UpdatedAt)
} }
@@ -155,7 +172,7 @@ func (s *PackageStore) ListForUser(ctx context.Context, userID string) ([]store.
&up.ID, &up.Title, &up.Type, &up.Version, &up.Description, &up.ID, &up.Title, &up.Type, &up.Version, &up.Description,
&up.Author, &up.Tier, &up.IsSystem, &up.Scope, &up.Author, &up.Tier, &up.IsSystem, &up.Scope,
&teamID, &installedBy, &teamID, &installedBy,
&manifestJSON, &up.Enabled, &up.Source, &manifestJSON, &up.Enabled, &up.Status, &up.Source,
&up.InstalledAt, &up.UpdatedAt, &up.InstalledAt, &up.UpdatedAt,
&userEnabled, &userSettings, &userEnabled, &userSettings,
); err != nil { ); err != nil {
@@ -218,7 +235,7 @@ func (s *PackageStore) DeleteUserSettings(ctx context.Context, pkgID, userID str
// so column additions don't silently break positional Scan(). // so column additions don't silently break positional Scan().
const pkgCols = `p.id, p.title, p.type, p.version, p.description, p.author, const pkgCols = `p.id, p.title, p.type, p.version, p.description, p.author,
p.tier, p.is_system, p.scope, p.team_id, p.installed_by, p.tier, p.is_system, p.scope, p.team_id, p.installed_by,
p.manifest, p.enabled, p.source, p.installed_at, p.updated_at` p.manifest, p.enabled, p.status, p.source, p.installed_at, p.updated_at`
func (s *PackageStore) scanOne(ctx context.Context, query string, args ...interface{}) (*store.PackageRegistration, error) { func (s *PackageStore) scanOne(ctx context.Context, query string, args ...interface{}) (*store.PackageRegistration, error) {
var pkg store.PackageRegistration var pkg store.PackageRegistration
@@ -228,7 +245,7 @@ func (s *PackageStore) scanOne(ctx context.Context, query string, args ...interf
&pkg.ID, &pkg.Title, &pkg.Type, &pkg.Version, &pkg.Description, &pkg.ID, &pkg.Title, &pkg.Type, &pkg.Version, &pkg.Description,
&pkg.Author, &pkg.Tier, &pkg.IsSystem, &pkg.Scope, &pkg.Author, &pkg.Tier, &pkg.IsSystem, &pkg.Scope,
&teamID, &installedBy, &teamID, &installedBy,
&manifestJSON, &pkg.Enabled, &pkg.Source, &manifestJSON, &pkg.Enabled, &pkg.Status, &pkg.Source,
&pkg.InstalledAt, &pkg.UpdatedAt, &pkg.InstalledAt, &pkg.UpdatedAt,
) )
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
@@ -259,7 +276,7 @@ func (s *PackageStore) scanMany(ctx context.Context, query string, args ...inter
&pkg.ID, &pkg.Title, &pkg.Type, &pkg.Version, &pkg.Description, &pkg.ID, &pkg.Title, &pkg.Type, &pkg.Version, &pkg.Description,
&pkg.Author, &pkg.Tier, &pkg.IsSystem, &pkg.Scope, &pkg.Author, &pkg.Tier, &pkg.IsSystem, &pkg.Scope,
&teamID, &installedBy, &teamID, &installedBy,
&manifestJSON, &pkg.Enabled, &pkg.Source, &manifestJSON, &pkg.Enabled, &pkg.Status, &pkg.Source,
&pkg.InstalledAt, &pkg.UpdatedAt, &pkg.InstalledAt, &pkg.UpdatedAt,
); err != nil { ); err != nil {
return nil, err return nil, err

View File

@@ -0,0 +1,138 @@
package postgres
import (
"context"
"database/sql"
"git.gobha.me/xcaliber/chat-switchboard/models"
)
type PersonaGroupStore struct{}
func NewPersonaGroupStore() *PersonaGroupStore { return &PersonaGroupStore{} }
const personaGroupCols = `id, name, description, owner_id, scope, team_id, created_at, updated_at`
func (s *PersonaGroupStore) List(ctx context.Context, ownerID string) ([]models.PersonaGroup, error) {
rows, err := DB.QueryContext(ctx, `
SELECT `+personaGroupCols+` FROM persona_groups
WHERE owner_id = $1 ORDER BY name
`, ownerID)
if err != nil {
return nil, err
}
defer rows.Close()
groups := []models.PersonaGroup{}
for rows.Next() {
var g models.PersonaGroup
if err := rows.Scan(&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, &g.CreatedAt, &g.UpdatedAt); err != nil {
continue
}
groups = append(groups, g)
}
return groups, rows.Err()
}
func (s *PersonaGroupStore) Get(ctx context.Context, id, ownerID string) (*models.PersonaGroup, error) {
var g models.PersonaGroup
err := DB.QueryRowContext(ctx, `
SELECT `+personaGroupCols+` FROM persona_groups WHERE id = $1 AND owner_id = $2
`, id, ownerID).Scan(&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, &g.CreatedAt, &g.UpdatedAt)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &g, nil
}
func (s *PersonaGroupStore) Create(ctx context.Context, g *models.PersonaGroup) error {
return DB.QueryRowContext(ctx, `
INSERT INTO persona_groups (name, description, owner_id, scope)
VALUES ($1, $2, $3, $4)
RETURNING id, name, description, owner_id, scope, team_id, created_at, updated_at
`, g.Name, g.Description, g.OwnerID, g.Scope).Scan(
&g.ID, &g.Name, &g.Description, &g.OwnerID,
&g.Scope, &g.TeamID, &g.CreatedAt, &g.UpdatedAt)
}
func (s *PersonaGroupStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
for k, v := range fields {
_, err := DB.ExecContext(ctx,
`UPDATE persona_groups SET `+k+` = $1, updated_at = NOW() WHERE id = $2`, v, id)
if err != nil {
return err
}
}
return nil
}
func (s *PersonaGroupStore) Delete(ctx context.Context, id, ownerID string) (int64, error) {
result, err := DB.ExecContext(ctx,
`DELETE FROM persona_groups WHERE id = $1 AND owner_id = $2`, id, ownerID)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *PersonaGroupStore) GetOwnerID(ctx context.Context, id string) (string, error) {
var ownerID string
err := DB.QueryRowContext(ctx,
`SELECT owner_id FROM persona_groups WHERE id = $1`, id).Scan(&ownerID)
if err == sql.ErrNoRows {
return "", nil
}
return ownerID, err
}
func (s *PersonaGroupStore) AddMember(ctx context.Context, groupID, personaID string, isLeader bool) error {
if isLeader {
_, _ = DB.ExecContext(ctx,
`UPDATE persona_group_members SET is_leader = false WHERE group_id = $1`, groupID)
}
_, err := DB.ExecContext(ctx, `
INSERT INTO persona_group_members (group_id, persona_id, is_leader)
VALUES ($1, $2, $3)
ON CONFLICT (group_id, persona_id) DO UPDATE SET is_leader = EXCLUDED.is_leader
`, groupID, personaID, isLeader)
return err
}
func (s *PersonaGroupStore) RemoveMember(ctx context.Context, memberID, groupID string) error {
_, err := DB.ExecContext(ctx,
`DELETE FROM persona_group_members WHERE id = $1 AND group_id = $2`, memberID, groupID)
return err
}
func (s *PersonaGroupStore) ListMembers(ctx context.Context, groupID string) ([]models.PersonaGroupMember, error) {
rows, err := DB.QueryContext(ctx, `
SELECT pgm.id, pgm.group_id, pgm.persona_id, pgm.is_leader, pgm.sort_order,
COALESCE(p.name, '') AS persona_name,
COALESCE(p.handle, '') AS persona_handle,
COALESCE(p.avatar, '') AS persona_avatar
FROM persona_group_members pgm
LEFT JOIN personas p ON p.id = pgm.persona_id
WHERE pgm.group_id = $1
ORDER BY pgm.is_leader DESC, pgm.sort_order, pgm.id
`, groupID)
if err != nil {
return []models.PersonaGroupMember{}, nil
}
defer rows.Close()
members := []models.PersonaGroupMember{}
for rows.Next() {
var m models.PersonaGroupMember
if err := rows.Scan(&m.ID, &m.GroupID, &m.PersonaID, &m.IsLeader, &m.SortOrder,
&m.PersonaName, &m.PersonaHandle, &m.PersonaAvatar); err != nil {
continue
}
members = append(members, m)
}
return members, rows.Err()
}

View File

@@ -0,0 +1,80 @@
package postgres
import (
"context"
"database/sql"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ── Mention resolution + display info (v0.29.0) ────────────────────────
func (s *PersonaStore) FindActiveByHandle(ctx context.Context, handle string) (string, error) {
var id string
err := DB.QueryRowContext(ctx, `
SELECT id FROM personas
WHERE LOWER(handle) = LOWER($1) AND is_active = true
LIMIT 1
`, handle).Scan(&id)
if err == sql.ErrNoRows {
return "", nil
}
return id, err
}
func (s *PersonaStore) FindActiveByHandlePrefix(ctx context.Context, prefix string) (string, int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM personas WHERE LOWER(handle) LIKE LOWER($1) AND is_active = true
`, prefix+"%").Scan(&count)
if err != nil {
return "", 0, err
}
if count != 1 {
return "", count, nil
}
var id string
err = DB.QueryRowContext(ctx, `
SELECT id FROM personas WHERE LOWER(handle) LIKE LOWER($1) AND is_active = true
`, prefix+"%").Scan(&id)
return id, 1, err
}
func (s *PersonaStore) GetNameByID(ctx context.Context, id string) (string, error) {
var name string
err := DB.QueryRowContext(ctx, `SELECT name FROM personas WHERE id = $1`, id).Scan(&name)
if err == sql.ErrNoRows {
return "", nil
}
return name, err
}
func (s *PersonaStore) GetNamesByIDs(ctx context.Context, ids []string) (map[string]string, error) {
result := make(map[string]string)
for _, id := range ids {
var name string
err := DB.QueryRowContext(ctx, `SELECT name FROM personas WHERE id = $1`, id).Scan(&name)
if err == nil && name != "" {
result[id] = name
}
}
return result, nil
}
func (s *PersonaStore) GetDisplayInfoByIDs(ctx context.Context, ids []string) (map[string]store.UserDisplayInfo, error) {
result := make(map[string]store.UserDisplayInfo)
for _, id := range ids {
var name, avatar sql.NullString
_ = DB.QueryRowContext(ctx, `
SELECT name, avatar FROM personas WHERE id = $1
`, id).Scan(&name, &avatar)
if name.Valid {
info := store.UserDisplayInfo{Name: name.String}
if avatar.Valid {
info.Avatar = avatar.String
}
result[id] = info
}
}
return result, nil
}

View File

@@ -0,0 +1,50 @@
package postgres
import (
"context"
"database/sql"
"time"
)
// PresenceStore manages user_presence table.
type PresenceStore struct{}
func NewPresenceStore() *PresenceStore { return &PresenceStore{} }
func (s *PresenceStore) Heartbeat(ctx context.Context, userID string) error {
_, err := DB.ExecContext(ctx, `
INSERT INTO user_presence (user_id, last_seen, status)
VALUES ($1, NOW(), 'online')
ON CONFLICT (user_id) DO UPDATE
SET last_seen = NOW(), status = 'online'
`, userID)
return err
}
func (s *PresenceStore) GetLastSeen(ctx context.Context, userID string) (*time.Time, error) {
var lastSeen time.Time
err := DB.QueryRowContext(ctx,
`SELECT last_seen FROM user_presence WHERE user_id = $1`, userID).Scan(&lastSeen)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
return &lastSeen, nil
}
func (s *PresenceStore) GetStatuses(ctx context.Context, userIDs []string, threshold time.Time) (map[string]string, error) {
result := make(map[string]string, len(userIDs))
for _, id := range userIDs {
var lastSeen time.Time
err := DB.QueryRowContext(ctx,
`SELECT last_seen FROM user_presence WHERE user_id = $1`, id).Scan(&lastSeen)
if err != nil || lastSeen.Before(threshold) {
result[id] = "offline"
} else {
result[id] = "online"
}
}
return result, nil
}

View File

@@ -8,6 +8,7 @@ import (
"github.com/lib/pq" "github.com/lib/pq"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
// ── ProjectStore ─────────────────────────── // ── ProjectStore ───────────────────────────
@@ -395,6 +396,45 @@ func (s *ProjectStore) GetProjectIDForChannel(ctx context.Context, channelID str
return projectID, err return projectID, err
} }
// ── CS5c additions (v0.29.0) ──────────────────────────────────────────
func (s *ProjectStore) AdminList(ctx context.Context, includeArchived bool) ([]store.AdminProject, error) {
rows, err := DB.QueryContext(ctx, `
SELECT p.id, p.name, p.description, p.scope,
p.owner_id, p.team_id, p.is_archived,
p.created_at, p.updated_at,
(SELECT COUNT(*) FROM project_channels WHERE project_id = p.id),
(SELECT COUNT(*) FROM project_knowledge_bases WHERE project_id = p.id),
(SELECT COUNT(*) FROM project_notes WHERE project_id = p.id),
COALESCE(u.username, '')
FROM projects p
LEFT JOIN users u ON u.id = p.owner_id
WHERE ($1 OR p.is_archived = false)
ORDER BY p.updated_at DESC`, includeArchived)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.AdminProject, 0)
for rows.Next() {
var p store.AdminProject
var teamID sql.NullString
if err := rows.Scan(
&p.ID, &p.Name, &p.Description, &p.Scope,
&p.OwnerID, &teamID, &p.IsArchived,
&p.CreatedAt, &p.UpdatedAt,
&p.ChannelCount, &p.KBCount, &p.NoteCount,
&p.OwnerName,
); err != nil {
return nil, err
}
p.TeamID = NullableStringPtr(teamID)
results = append(results, p)
}
return results, rows.Err()
}
// ── Query Helper ──────────────────────────── // ── Query Helper ────────────────────────────
func queryProjects(ctx context.Context, q string, args ...interface{}) ([]models.Project, error) { func queryProjects(ctx context.Context, q string, args ...interface{}) ([]models.Project, error) {

View File

@@ -195,3 +195,64 @@ func scanProviders(rows *sql.Rows) ([]models.ProviderConfig, error) {
} }
return result, rows.Err() return result, rows.Err()
} }
// ── CS4 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ProviderStore) DeletePersonalByOwner(ctx context.Context, ownerID string) (int64, error) {
result, err := DB.ExecContext(ctx,
`DELETE FROM provider_configs WHERE scope = 'personal' AND owner_id = $1`, ownerID)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
// ── CS6 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ProviderStore) ListAllForTeam(ctx context.Context, teamID string) ([]models.ProviderConfig, error) {
rows, err := DB.QueryContext(ctx,
fmt.Sprintf("SELECT %s FROM provider_configs WHERE scope = 'team' AND owner_id = $1 ORDER BY name", providerCols),
teamID)
if err != nil {
return nil, err
}
defer rows.Close()
return scanProviders(rows)
}
func (s *ProviderStore) DeleteByIDAndTeam(ctx context.Context, id, teamID string) (int64, error) {
res, err := DB.ExecContext(ctx,
`DELETE FROM provider_configs WHERE id = $1 AND scope = 'team' AND owner_id = $2`,
id, teamID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// ── CS7a additions (v0.29.0) ────────────────────────────────────────────
func (s *ProviderStore) FindFirstForUser(ctx context.Context, userID string) (string, error) {
var configID string
err := DB.QueryRowContext(ctx, `
SELECT id FROM provider_configs
WHERE is_active = true AND (
(scope = 'personal' AND owner_id = $1)
OR scope = 'global'
)
ORDER BY scope ASC, created_at ASC
LIMIT 1
`, userID).Scan(&configID)
return configID, err
}
func (s *ProviderStore) LoadAccessible(ctx context.Context, configID, userID string) (*models.ProviderConfig, error) {
row := DB.QueryRowContext(ctx, fmt.Sprintf(`
SELECT %s FROM provider_configs
WHERE id = $1 AND is_active = true
AND (scope = 'global'
OR (scope = 'personal' AND owner_id = $2)
OR (scope = 'team' AND owner_id IN (SELECT team_id FROM team_members WHERE user_id = $2)))
`, providerCols), configID, userID)
return scanProvider(row)
}

View File

@@ -42,5 +42,10 @@ func NewStores(db *sql.DB) store.Stores {
Packages: NewPackageStore(), Packages: NewPackageStore(),
Workflows: NewWorkflowStore(), Workflows: NewWorkflowStore(),
Tasks: NewTaskStore(), Tasks: NewTaskStore(),
Presence: NewPresenceStore(),
PersonaGroups: NewPersonaGroupStore(),
Folders: NewFolderStore(),
Health: NewHealthStore(db),
ExtPermissions: NewExtensionPermissionStore(db),
} }
} }

View File

@@ -226,3 +226,101 @@ func (s *TeamStore) IsMember(ctx context.Context, teamID, userID string) (bool,
// unused but keeping for reference // unused but keeping for reference
var _ = fmt.Sprintf var _ = fmt.Sprintf
// ── CS1 additions (v0.29.0) ─────────────────────────────────────────────
func (s *TeamStore) Exists(ctx context.Context, teamID string) (bool, error) {
var exists bool
err := DB.QueryRowContext(ctx,
`SELECT EXISTS(SELECT 1 FROM teams WHERE id = $1)`, teamID).Scan(&exists)
return exists, err
}
func (s *TeamStore) UpdateMemberRoleByID(ctx context.Context, memberID, teamID, role string) (int64, error) {
res, err := DB.ExecContext(ctx,
`UPDATE team_members SET role = $1 WHERE id = $2 AND team_id = $3`,
role, memberID, teamID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *TeamStore) DeleteMemberByID(ctx context.Context, memberID, teamID string) (int64, error) {
res, err := DB.ExecContext(ctx,
`DELETE FROM team_members WHERE id = $1 AND team_id = $2`,
memberID, teamID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *TeamStore) ListTeamAuditActions(ctx context.Context, teamID string) ([]string, error) {
rows, err := DB.QueryContext(ctx, `
SELECT DISTINCT al.action
FROM audit_log al
WHERE al.actor_id IN (SELECT user_id FROM team_members WHERE team_id = $1)
ORDER BY al.action
`, teamID)
if err != nil {
return nil, err
}
defer rows.Close()
var actions []string
for rows.Next() {
var a string
if err := rows.Scan(&a); err != nil {
return nil, err
}
actions = append(actions, a)
}
if actions == nil {
actions = []string{}
}
return actions, rows.Err()
}
// ── CS5b additions (v0.29.0) ────────────────────────────────────────────
func (s *TeamStore) GetFirstTeamIDForUser(ctx context.Context, userID string) (string, error) {
var teamID string
err := DB.QueryRowContext(ctx,
`SELECT team_id FROM team_members WHERE user_id = $1 LIMIT 1`, userID).Scan(&teamID)
if err != nil {
return "", nil
}
return teamID, nil
}
// ── CS6 additions (v0.29.0) ────────────────────────────────────────────
func (s *TeamStore) AddMemberReturningID(ctx context.Context, teamID, userID, role string) (string, error) {
var id string
err := DB.QueryRowContext(ctx, `
INSERT INTO team_members (team_id, user_id, role)
VALUES ($1, $2, $3)
RETURNING id`, teamID, userID, role).Scan(&id)
return id, err
}
func (s *TeamStore) HasPrivateProviderRequirement(ctx context.Context, userID string) (bool, error) {
var has bool
err := DB.QueryRowContext(ctx, `
SELECT EXISTS(
SELECT 1 FROM team_members tm
JOIN teams t ON t.id = tm.team_id
WHERE tm.user_id = $1
AND t.is_active = true
AND t.settings->>'require_private_providers' = 'true'
)`, userID).Scan(&has)
return has, err
}
func (s *TeamStore) MergeSettings(ctx context.Context, teamID, settingsJSON string) error {
_, err := DB.ExecContext(ctx,
`UPDATE teams SET settings = COALESCE(settings, '{}'::jsonb) || $1::jsonb WHERE id = $2`,
settingsJSON, teamID)
return err
}

View File

@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"fmt" "fmt"
"strings"
"time" "time"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
@@ -186,3 +187,110 @@ func scanOneUser(ctx context.Context, query string, args ...interface{}) (*model
ScanJSON(sj, &u.Settings) ScanJSON(sj, &u.Settings)
return &u, nil return &u, nil
} }
// ── CS1 additions (v0.29.0) ─────────────────────────────────────────────
func (s *UserStore) Exists(ctx context.Context, userID string) (bool, error) {
var exists bool
err := DB.QueryRowContext(ctx,
`SELECT EXISTS(SELECT 1 FROM users WHERE id = $1)`, userID).Scan(&exists)
return exists, err
}
func (s *UserStore) SearchActive(ctx context.Context, excludeUserID, query string) ([]store.UserSearchResult, error) {
q := `
SELECT id, username, COALESCE(display_name, '') AS display_name, COALESCE(handle, '') AS handle
FROM users
WHERE is_active = true AND id != $1`
args := []interface{}{excludeUserID}
if query != "" {
q += ` AND (LOWER(username) LIKE $2 OR LOWER(display_name) LIKE $3 OR LOWER(handle) LIKE $4)`
pattern := "%" + strings.ToLower(query) + "%"
args = append(args, pattern, pattern, pattern)
}
q += ` ORDER BY username LIMIT 20`
rows, err := DB.QueryContext(ctx, q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var results []store.UserSearchResult
for rows.Next() {
var u store.UserSearchResult
if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Handle); err != nil {
continue
}
results = append(results, u)
}
if results == nil {
results = []store.UserSearchResult{}
}
return results, rows.Err()
}
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *UserStore) CountByRole(ctx context.Context, role string) (int, error) {
var count int
err := DB.QueryRowContext(ctx,
`SELECT COUNT(*) FROM users WHERE role = $1`, role).Scan(&count)
return count, err
}
func (s *UserStore) MergeSettings(ctx context.Context, userID string, patch []byte) error {
_, err := DB.ExecContext(ctx, `
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
`, string(patch), userID)
return err
}
func (s *UserStore) GetVaultKeys(ctx context.Context, userID string) (bool, []byte, []byte, []byte, error) {
var vaultSet bool
var encUEK, salt, nonce []byte
err := DB.QueryRowContext(ctx, `
SELECT vault_set, encrypted_uek, uek_salt, uek_nonce
FROM users WHERE id = $1
`, userID).Scan(&vaultSet, &encUEK, &salt, &nonce)
return vaultSet, encUEK, salt, nonce, err
}
func (s *UserStore) UpdateVaultKeys(ctx context.Context, userID string, encUEK, salt, nonce []byte) error {
_, err := DB.ExecContext(ctx, `
UPDATE users
SET encrypted_uek = $1, uek_salt = $2, uek_nonce = $3, updated_at = NOW()
WHERE id = $4
`, encUEK, salt, nonce, userID)
return err
}
func (s *UserStore) CountAll(ctx context.Context) (int, error) {
var count int
err := DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&count)
return count, err
}
// ── CS4 additions (v0.29.0) ─────────────────────────────────────────────
func (s *UserStore) ClearVaultKeys(ctx context.Context, userID string) error {
_, err := DB.ExecContext(ctx, `
UPDATE users
SET encrypted_uek = NULL, uek_salt = NULL, uek_nonce = NULL, vault_set = false
WHERE id = $1
`, userID)
return err
}
func (s *UserStore) InitVaultKeys(ctx context.Context, userID string, encUEK, salt, nonce []byte) error {
_, err := DB.ExecContext(ctx, `
UPDATE users
SET encrypted_uek = $1, uek_salt = $2, uek_nonce = $3, vault_set = true
WHERE id = $4
`, encUEK, salt, nonce, userID)
return err
}

View File

@@ -0,0 +1,62 @@
package postgres
import (
"context"
"database/sql"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
// ── Mention resolution + display info (v0.29.0) ────────────────────────
func (s *UserStore) FindActiveByHandle(ctx context.Context, handle, excludeUserID string) (string, error) {
var id string
err := DB.QueryRowContext(ctx, `
SELECT id FROM users
WHERE LOWER(handle) = LOWER($1) AND id != $2 AND is_active = true
LIMIT 1
`, handle, excludeUserID).Scan(&id)
if err == sql.ErrNoRows {
return "", nil
}
return id, err
}
func (s *UserStore) FindActiveByHandlePrefix(ctx context.Context, prefix, excludeUserID string) (string, int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM users
WHERE LOWER(handle) LIKE LOWER($1) AND id != $2 AND is_active = true
`, prefix+"%", excludeUserID).Scan(&count)
if err != nil {
return "", 0, err
}
if count != 1 {
return "", count, nil
}
var id string
err = DB.QueryRowContext(ctx, `
SELECT id FROM users
WHERE LOWER(handle) LIKE LOWER($1) AND id != $2 AND is_active = true
LIMIT 1
`, prefix+"%", excludeUserID).Scan(&id)
return id, 1, err
}
func (s *UserStore) GetDisplayInfoByIDs(ctx context.Context, ids []string) (map[string]store.UserDisplayInfo, error) {
result := make(map[string]store.UserDisplayInfo)
for _, id := range ids {
var name, avatar sql.NullString
_ = DB.QueryRowContext(ctx, `
SELECT COALESCE(display_name, username), avatar_url FROM users WHERE id = $1
`, id).Scan(&name, &avatar)
if name.Valid {
info := store.UserDisplayInfo{Name: name.String}
if avatar.Valid {
info.Avatar = avatar.String
}
result[id] = info
}
}
return result, nil
}

View File

@@ -2,10 +2,13 @@ package postgres
import ( import (
"context" "context"
"database/sql"
"encoding/json" "encoding/json"
"fmt" "fmt"
"time"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
) )
// WorkflowStore implements store.WorkflowStore for Postgres. // WorkflowStore implements store.WorkflowStore for Postgres.
@@ -357,3 +360,129 @@ func nullIfEmpty(s string) interface{} {
} }
return s return s
} }
// ── Assignments (v0.29.0-cs3) ───────────────────────────────────────────
func (s *WorkflowStore) CreateAssignment(ctx context.Context, a *store.WorkflowAssignment) error {
a.ID = store.NewID()
_, err := DB.ExecContext(ctx, `
INSERT INTO workflow_assignments (id, channel_id, stage, team_id)
VALUES ($1, $2, $3, $4)
`, a.ID, a.ChannelID, a.Stage, a.TeamID)
return err
}
func (s *WorkflowStore) ListAssignmentsForTeam(ctx context.Context, teamID, status string) ([]store.WorkflowAssignment, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, channel_id, stage, team_id, assigned_to, status,
created_at, claimed_at, completed_at
FROM workflow_assignments
WHERE team_id = $1 AND status = $2
ORDER BY created_at DESC
`, teamID, status)
if err != nil {
return nil, err
}
defer rows.Close()
return scanAssignments(rows)
}
func (s *WorkflowStore) ListAssignmentsMine(ctx context.Context, userID string) ([]store.WorkflowAssignment, error) {
rows, err := DB.QueryContext(ctx, `
SELECT DISTINCT wa.id, wa.channel_id, wa.stage, wa.team_id, wa.assigned_to, wa.status,
wa.created_at, wa.claimed_at, wa.completed_at
FROM workflow_assignments wa
LEFT JOIN team_members tm ON tm.team_id = wa.team_id AND tm.user_id = $1
WHERE (wa.assigned_to = $2 AND wa.status = 'claimed')
OR (wa.status = 'unassigned' AND tm.user_id IS NOT NULL)
ORDER BY wa.created_at DESC
`, userID, userID)
if err != nil {
return nil, err
}
defer rows.Close()
return scanAssignments(rows)
}
func (s *WorkflowStore) ClaimAssignment(ctx context.Context, assignmentID, userID string) (int64, error) {
res, err := DB.ExecContext(ctx, `
UPDATE workflow_assignments
SET assigned_to = $1, status = 'claimed', claimed_at = $2
WHERE id = $3 AND status = 'unassigned'
`, userID, time.Now().UTC(), assignmentID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *WorkflowStore) CompleteAssignment(ctx context.Context, assignmentID string) (int64, error) {
res, err := DB.ExecContext(ctx, `
UPDATE workflow_assignments
SET status = 'completed', completed_at = $1
WHERE id = $2 AND status = 'claimed'
`, time.Now().UTC(), assignmentID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *WorkflowStore) GetAssignmentChannelID(ctx context.Context, assignmentID string) (string, error) {
var channelID string
err := DB.QueryRowContext(ctx,
`SELECT channel_id FROM workflow_assignments WHERE id = $1`, assignmentID).Scan(&channelID)
return channelID, err
}
func (s *WorkflowStore) TryRoundRobin(ctx context.Context, teamID, assignmentID string) (string, error) {
// Find least-recently-assigned team member
rows, err := DB.QueryContext(ctx, `
SELECT m.user_id, COALESCE(MAX(wa.claimed_at), '1970-01-01T00:00:00Z') as last_claim
FROM team_members m
LEFT JOIN workflow_assignments wa ON wa.assigned_to = m.user_id AND wa.team_id = $1
WHERE m.team_id = $2
GROUP BY m.user_id
ORDER BY last_claim ASC
LIMIT 1
`, teamID, teamID)
if err != nil {
return "", err
}
defer rows.Close()
if !rows.Next() {
return "", nil // no team members
}
var userID, lastClaim string
if err := rows.Scan(&userID, &lastClaim); err != nil {
return "", err
}
// Claim for that user
_, err = DB.ExecContext(ctx, `
UPDATE workflow_assignments
SET assigned_to = $1, status = 'claimed', claimed_at = $2
WHERE id = $3 AND status = 'unassigned'
`, userID, time.Now().UTC(), assignmentID)
if err != nil {
return "", err
}
return userID, nil
}
func scanAssignments(rows *sql.Rows) ([]store.WorkflowAssignment, error) {
var result []store.WorkflowAssignment
for rows.Next() {
var a store.WorkflowAssignment
if err := rows.Scan(&a.ID, &a.ChannelID, &a.Stage, &a.TeamID,
&a.AssignedTo, &a.Status, &a.CreatedAt, &a.ClaimedAt, &a.CompletedAt); err != nil {
return nil, err
}
result = append(result, a)
}
if result == nil {
result = []store.WorkflowAssignment{}
}
return result, rows.Err()
}

View File

@@ -2,6 +2,7 @@ package store
import ( import (
"context" "context"
"time"
"git.gobha.me/xcaliber/chat-switchboard/models" "git.gobha.me/xcaliber/chat-switchboard/models"
) )
@@ -44,4 +45,26 @@ type ProjectStore interface {
// Get KB IDs bound to a project (used for virtual injection at completion) // Get KB IDs bound to a project (used for virtual injection at completion)
GetKBIDs(ctx context.Context, projectID string) ([]string, error) GetKBIDs(ctx context.Context, projectID string) ([]string, error)
// ── CS5c additions (v0.29.0) ──
// AdminList returns all projects with counts and owner names (admin only).
AdminList(ctx context.Context, includeArchived bool) ([]AdminProject, error)
}
// AdminProject is returned by ProjectStore.AdminList.
type AdminProject struct {
ID string
Name string
Description string
Scope string
OwnerID string
TeamID *string
IsArchived bool
CreatedAt time.Time
UpdatedAt time.Time
ChannelCount int
KBCount int
NoteCount int
OwnerName string
} }

View File

@@ -34,6 +34,9 @@ func (s *AuditStore) List(ctx context.Context, opts store.AuditListOptions) ([]m
al.ip_address, al.user_agent, al.created_at`, "audit_log al") al.ip_address, al.user_agent, al.created_at`, "audit_log al")
b.Join("LEFT JOIN users u ON al.actor_id = u.id") b.Join("LEFT JOIN users u ON al.actor_id = u.id")
if opts.TeamID != "" {
b.Where("al.actor_id IN (SELECT user_id FROM team_members WHERE team_id = ?)", opts.TeamID)
}
if opts.ActorID != "" { if opts.ActorID != "" {
b.Where("al.actor_id = ?", opts.ActorID) b.Where("al.actor_id = ?", opts.ActorID)
} }

View File

@@ -276,3 +276,58 @@ func scanCatalogEntries(rows *sql.Rows) ([]models.CatalogEntry, error) {
} }
return result, rows.Err() return result, rows.Err()
} }
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *CatalogStore) GetCapabilities(ctx context.Context, modelID, configID string) ([]byte, error) {
var capsJSON []byte
err := DB.QueryRowContext(ctx, `
SELECT capabilities FROM model_catalog
WHERE model_id = ? AND provider_config_id = ?
`, modelID, configID).Scan(&capsJSON)
if err != nil {
return nil, err
}
return capsJSON, nil
}
func (s *CatalogStore) GetCapabilitiesAny(ctx context.Context, modelID string) ([]byte, error) {
var capsJSON []byte
err := DB.QueryRowContext(ctx, `
SELECT capabilities FROM model_catalog
WHERE model_id = ? ORDER BY last_synced_at DESC LIMIT 1
`, modelID).Scan(&capsJSON)
if err != nil {
return nil, err
}
return capsJSON, nil
}
// ── CS6 additions (v0.29.0) ─────────────────────────────────────────────
func (s *CatalogStore) ListTeamAvailable(ctx context.Context) ([]store.TeamAvailableModel, error) {
rows, err := DB.QueryContext(ctx, `
SELECT mc.id, mc.model_id, mc.display_name, mc.visibility,
ac.provider, ac.name AS provider_name
FROM model_catalog mc
JOIN provider_configs ac ON mc.provider_config_id = ac.id
WHERE mc.visibility IN ('enabled', 'team')
AND ac.is_active = 1 AND ac.scope = 'global'
ORDER BY ac.name, mc.model_id
`)
if err != nil {
return nil, err
}
defer rows.Close()
results := make([]store.TeamAvailableModel, 0)
for rows.Next() {
var m store.TeamAvailableModel
if err := rows.Scan(&m.ID, &m.ModelID, &m.DisplayName, &m.Visibility,
&m.Provider, &m.ProviderName); err != nil {
continue
}
results = append(results, m)
}
return results, rows.Err()
}

View File

@@ -0,0 +1,57 @@
package sqlite
import (
"context"
"database/sql"
)
// ── Mention resolution (v0.29.0) ────────────────────────────────────────
func (s *CatalogStore) FindEnabledByModelID(ctx context.Context, modelID string) (string, string, error) {
var foundModelID, configID string
err := DB.QueryRowContext(ctx, `
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) = LOWER(?)
AND mc.visibility = 'enabled'
AND pc.is_active = 1
ORDER BY
CASE pc.scope WHEN 'global' THEN 0 WHEN 'team' THEN 1 WHEN 'personal' THEN 2 END
LIMIT 1
`, modelID).Scan(&foundModelID, &configID)
if err == sql.ErrNoRows {
return "", "", nil
}
return foundModelID, configID, err
}
func (s *CatalogStore) FindEnabledByModelIDPrefix(ctx context.Context, prefix string) (string, string, int, error) {
var count int
err := DB.QueryRowContext(ctx, `
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 LOWER(?)
AND mc.visibility = 'enabled'
AND pc.is_active = 1
`, prefix+"%").Scan(&count)
if err != nil {
return "", "", 0, err
}
if count != 1 {
return "", "", count, nil
}
var foundModelID, configID string
err = DB.QueryRowContext(ctx, `
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 LOWER(?)
AND mc.visibility = 'enabled'
AND pc.is_active = 1
ORDER BY CASE pc.scope WHEN 'global' THEN 0 WHEN 'team' THEN 1 WHEN 'personal' THEN 2 END
LIMIT 1
`, prefix+"%").Scan(&foundModelID, &configID)
return foundModelID, configID, 1, err
}

View File

@@ -477,3 +477,427 @@ func (s *ChannelStore) Purge(ctx context.Context, id string) error {
_, err = DB.ExecContext(ctx, `DELETE FROM channels WHERE id = ?`, id) _, err = DB.ExecContext(ctx, `DELETE FROM channels WHERE id = ?`, id)
return err return err
} }
// ── CS1 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ChannelStore) FindExistingDM(ctx context.Context, userID1, userID2 string) (string, error) {
var channelID string
err := DB.QueryRowContext(ctx, `
SELECT cp1.channel_id FROM channel_participants cp1
JOIN channel_participants cp2 ON cp1.channel_id = cp2.channel_id
JOIN channels c ON c.id = cp1.channel_id
WHERE c.type = 'dm'
AND cp1.participant_type = 'user' AND cp1.participant_id = ?
AND cp2.participant_type = 'user' AND cp2.participant_id = ?
LIMIT 1
`, userID1, userID2).Scan(&channelID)
if err == sql.ErrNoRows {
return "", nil
}
return channelID, err
}
func (s *ChannelStore) GetUnreadCount(ctx context.Context, channelID, userID string) (int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages m
JOIN channel_participants cp ON cp.channel_id = m.channel_id
WHERE cp.channel_id = ?
AND cp.participant_type = 'user' AND cp.participant_id = ?
AND m.created_at > cp.last_read_at
`, channelID, userID).Scan(&count)
return count, err
}
func (s *ChannelStore) DeleteByOwner(ctx context.Context, channelID, userID string) (int64, error) {
result, err := DB.ExecContext(ctx,
`DELETE FROM channels WHERE id = ? AND user_id = ?`,
channelID, userID)
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *ChannelStore) MarkRead(ctx context.Context, channelID, userID string) error {
_, err := DB.ExecContext(ctx, `
UPDATE channel_participants
SET last_read_at = datetime('now')
WHERE channel_id = ? AND participant_type = 'user' AND participant_id = ?
`, channelID, userID)
if err != nil {
return nil
}
var latestMsgID *string
_ = DB.QueryRowContext(ctx, `
SELECT id FROM messages WHERE channel_id = ? ORDER BY created_at DESC LIMIT 1
`, channelID).Scan(&latestMsgID)
if latestMsgID != nil {
_, _ = DB.ExecContext(ctx, `
UPDATE channel_participants
SET last_read_message_id = ?
WHERE channel_id = ? AND participant_type = 'user' AND participant_id = ?
`, *latestMsgID, channelID, userID)
}
return nil
}
func (s *ChannelStore) CountParticipantsByType(ctx context.Context, channelID, pType string) (int, error) {
var count int
err := DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM channel_participants
WHERE channel_id = ? AND participant_type = ?
`, channelID, pType).Scan(&count)
return count, err
}
// ── CS2 additions (v0.29.0) ─────────────────────────────────────────────
func (s *ChannelStore) CountAll(ctx context.Context) (int, error) {
var count int
err := DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM channels").Scan(&count)
return count, err
}
// ── Workflow instance state (v0.29.0-cs3) ───────────────────────────────
func (s *ChannelStore) SetWorkflowInstance(ctx context.Context, channelID, workflowID string, version int, stageData json.RawMessage, status string) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET workflow_id = ?, workflow_version = ?, current_stage = 0,
stage_data = ?, workflow_status = ?, last_activity_at = ?
WHERE id = ?
`, workflowID, version, stageData, status, time.Now().UTC().Format(timeFmt), channelID)
return err
}
func (s *ChannelStore) GetWorkflowStatus(ctx context.Context, channelID string) (*store.WorkflowChannelStatus, error) {
var ws store.WorkflowChannelStatus
var stageData []byte
err := DB.QueryRowContext(ctx, `
SELECT workflow_id, workflow_version, current_stage,
COALESCE(stage_data, '{}'), COALESCE(workflow_status, 'active'),
last_activity_at
FROM channels WHERE id = ? AND type = 'workflow'
`, channelID).Scan(&ws.WorkflowID, &ws.WorkflowVersion,
&ws.CurrentStage, &stageData, &ws.Status, &ws.LastActivityAt)
if err == sql.ErrNoRows {
return nil, nil
}
if err != nil {
return nil, err
}
ws.StageData = stageData
return &ws, nil
}
func (s *ChannelStore) AdvanceWorkflowStage(ctx context.Context, channelID string, nextStage int, stageData json.RawMessage) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET current_stage = ?, stage_data = ?, last_activity_at = ?
WHERE id = ?
`, nextStage, stageData, time.Now().UTC().Format(timeFmt), channelID)
return err
}
func (s *ChannelStore) CompleteWorkflow(ctx context.Context, channelID string, finalStage int, stageData json.RawMessage) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels
SET current_stage = ?, workflow_status = 'completed',
stage_data = ?, last_activity_at = ?, ai_mode = 'off'
WHERE id = ?
`, finalStage, stageData, time.Now().UTC().Format(timeFmt), channelID)
return err
}
func (s *ChannelStore) RejectWorkflowToStage(ctx context.Context, channelID string, stage int) error {
_, err := DB.ExecContext(ctx, `
UPDATE channels SET current_stage = ?, last_activity_at = ? WHERE id = ?
`, stage, time.Now().UTC().Format(timeFmt), channelID)
return err
}
func (s *ChannelStore) GetStageData(ctx context.Context, channelID string) (json.RawMessage, error) {
var data json.RawMessage
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(stage_data, '{}') FROM channels WHERE id = ?
`, channelID).Scan(&data)
return data, err
}
// ── Background job helpers (v0.29.0-cs4) ────────────────────────────────
func (s *ChannelStore) MarkStaleWorkflows(ctx context.Context, cutoff time.Time) (int64, error) {
result, err := DB.ExecContext(ctx, `
UPDATE channels
SET workflow_status = 'stale'
WHERE type = 'workflow'
AND workflow_status = 'active'
AND last_activity_at < ?
`, cutoff.Format(timeFmt))
if err != nil {
return 0, err
}
return result.RowsAffected()
}
func (s *ChannelStore) EnforceWorkflowRetention(ctx context.Context) (int64, error) {
// SQLite doesn't support JSON operators for retention policy evaluation.
// This is a no-op on SQLite; retention enforcement is PG-only.
return 0, nil
}
func (s *ChannelStore) GetTypeAndAllowAnonymous(ctx context.Context, channelID string) (string, bool, error) {
var chType string
var allowAnon bool
err := DB.QueryRowContext(ctx,
`SELECT type, allow_anonymous FROM channels WHERE id = ?`, channelID).Scan(&chType, &allowAnon)
return chType, allowAnon, err
}
// ── CS5b additions (v0.29.0) ────────────────────────────────────────────
func (s *ChannelStore) FindCompactionCandidates(ctx context.Context, activityBefore, createdAfter time.Time, minMessages, minChars, limit int) ([]models.Channel, error) {
// Auto-compaction scanner is PG-only (uses settings::text cast).
return []models.Channel{}, nil
}
// ── CS7a additions (v0.29.0) ────────────────────────────────────────────
func (s *ChannelStore) GetProviderConfigID(ctx context.Context, channelID string) (*string, error) {
var configID sql.NullString
err := DB.QueryRowContext(ctx,
`SELECT provider_config_id FROM channels WHERE id = ?`, channelID).Scan(&configID)
if err != nil {
return nil, err
}
return NullableStringPtr(configID), nil
}
func (s *ChannelStore) UserCanAccess(ctx context.Context, channelID, userID string) (bool, error) {
var ok bool
err := DB.QueryRowContext(ctx, `
SELECT EXISTS(
SELECT 1 FROM channels WHERE id = ? AND user_id = ?
UNION ALL
SELECT 1 FROM channel_participants
WHERE channel_id = ? AND participant_type = 'user' AND participant_id = ?
LIMIT 1
)`, channelID, userID, channelID, userID).Scan(&ok)
return ok, err
}
// ── CS7b additions (v0.29.0) ────────────────────────────────────────────
const channelListCols = `c.id, c.user_id, c.title, c.type, c.ai_mode, c.topic,
c.description, c.model, c.provider_config_id,
c.system_prompt, c.is_archived, c.is_pinned, c.folder, c.folder_id, c.project_id, c.workspace_id,
c.tags, c.settings,
COALESCE(mc.cnt, 0) AS message_count,
c.created_at, c.updated_at`
const channelListFrom = `channels c
LEFT JOIN (
SELECT channel_id, COUNT(*) AS cnt FROM messages WHERE deleted_at IS NULL GROUP BY channel_id
) mc ON mc.channel_id = c.id`
func (s *ChannelStore) ListFiltered(ctx context.Context, userID string, f store.ChannelListFilter) ([]store.ChannelListItem, int, error) {
b := NewSelect(channelListCols, channelListFrom)
b.Where("(c.user_id = ? OR c.id IN (SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = ?))", userID, userID)
archivedVal := 0
if f.Archived {
archivedVal = 1
}
b.Where("c.is_archived = ?", archivedVal)
if len(f.Types) == 1 {
b.Where("c.type = ?", f.Types[0])
} else if len(f.Types) > 1 {
placeholders := make([]string, len(f.Types))
for i, t := range f.Types {
placeholders[i] = "?"
b.args = append(b.args, t)
}
b.WhereRaw("c.type IN (" + strings.Join(placeholders, ",") + ")")
}
if f.Folder != "" {
b.Where("c.folder = ?", f.Folder)
}
if f.FolderID != "" {
b.Where("c.folder_id = ?", f.FolderID)
}
if f.Search != "" {
b.Where("c.title LIKE ?", "%"+f.Search+"%")
}
if f.ProjectID == "none" {
b.WhereRaw("c.project_id IS NULL")
} else if f.ProjectID != "" {
b.Where("c.project_id = ?", f.ProjectID)
}
// Count
countQ, countArgs := b.CountBuild()
var total int
DB.QueryRowContext(ctx, countQ, countArgs...).Scan(&total)
b.OrderBy("c.is_pinned DESC, c.updated_at", "DESC")
b.Paginate(f.ListOptions)
q, args := b.Build()
rows, err := DB.QueryContext(ctx, q, args...)
if err != nil {
return nil, 0, err
}
defer rows.Close()
items, err := scanChannelListItems(rows)
if err != nil {
return nil, 0, err
}
// Compute unread counts
for i := range items {
DB.QueryRowContext(ctx, `
SELECT COUNT(*) FROM messages m
JOIN channel_participants cp ON cp.channel_id = m.channel_id
WHERE cp.channel_id = ?
AND cp.participant_type = 'user' AND cp.participant_id = ?
AND m.created_at > cp.last_read_at
AND m.deleted_at IS NULL
`, items[i].ID, userID).Scan(&items[i].UnreadCount)
}
return items, total, nil
}
func (s *ChannelStore) GetForUser(ctx context.Context, channelID, userID string) (*store.ChannelListItem, error) {
row := DB.QueryRowContext(ctx, fmt.Sprintf(`
SELECT %s FROM %s
WHERE c.id = ? AND (c.user_id = ? OR c.id IN (
SELECT channel_id FROM channel_participants WHERE participant_type = 'user' AND participant_id = ?
))
`, channelListCols, channelListFrom), channelID, userID, userID)
var item store.ChannelListItem
var tags, settings string
err := row.Scan(
&item.ID, &item.UserID, &item.Title, &item.Type, &item.AiMode, &item.Topic,
&item.Description, &item.Model, &item.ProviderConfigID,
&item.SystemPrompt, &item.IsArchived, &item.IsPinned, &item.Folder, &item.FolderID, &item.ProjectID, &item.WorkspaceID,
&tags, &settings,
&item.MessageCount, st(&item.CreatedAtTime), st(&item.UpdatedAtTime),
)
if err != nil {
return nil, err
}
item.Tags = ScanArray(tags)
if item.Tags == nil {
item.Tags = []string{}
}
item.Settings = safeJSONString(settings)
item.CreatedAt = item.CreatedAtTime.Format("2006-01-02T15:04:05Z")
item.UpdatedAt = item.UpdatedAtTime.Format("2006-01-02T15:04:05Z")
return &item, nil
}
func (s *ChannelStore) CreateFull(ctx context.Context, ch *models.Channel, folder string, tags []string, aiMode string,
ownerUserID string, dmPartnerIDs []string, defaultModel, defaultConfigID string) error {
ch.ID = store.NewID()
now := time.Now().UTC()
ch.CreatedAt = now
ch.UpdatedAt = now
if tags == nil {
tags = []string{}
}
tx, err := DB.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
tagsJSON := ArrayToJSON(tags)
_, err = tx.ExecContext(ctx, `
INSERT INTO channels (id, user_id, title, type, description, model, system_prompt,
provider_config_id, folder, folder_id, tags, ai_mode, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
ch.ID, ch.UserID, ch.Title, ch.Type, ch.Description, ch.Model, ch.SystemPrompt,
models.NullString(ch.ProviderConfigID), folder, models.NullString(ch.FolderID),
tagsJSON, aiMode, now.Format(timeFmt), now.Format(timeFmt),
)
if err != nil {
return fmt.Errorf("CreateFull insert channel: %w", err)
}
// Add owner participant
_, _ = tx.ExecContext(ctx, `
INSERT INTO channel_participants (id, channel_id, participant_type, participant_id, role)
VALUES (?, ?, 'user', ?, 'owner')
ON CONFLICT DO NOTHING`, store.NewID(), ch.ID, ownerUserID)
// Add DM partner participants
for _, pid := range dmPartnerIDs {
if pid == ownerUserID {
continue
}
_, _ = tx.ExecContext(ctx, `
INSERT INTO channel_participants (id, channel_id, participant_type, participant_id, role)
VALUES (?, ?, 'user', ?, 'member')
ON CONFLICT DO NOTHING`, store.NewID(), ch.ID, pid)
}
// Auto-create channel_model if model specified
if defaultModel != "" {
_, _ = tx.ExecContext(ctx, `
INSERT INTO channel_models (id, channel_id, model_id, provider_config_id, is_default)
VALUES (?, ?, ?, ?, 1)
ON CONFLICT DO NOTHING`, store.NewID(), ch.ID, defaultModel, models.NullString(&defaultConfigID))
}
return tx.Commit()
}
func (s *ChannelStore) MergeSettings(ctx context.Context, channelID string, settingsJSON json.RawMessage) error {
_, err := DB.ExecContext(ctx,
`UPDATE channels SET settings = json_patch(COALESCE(settings, '{}'), ?) WHERE id = ?`,
string(settingsJSON), channelID)
return err
}
// ── CS7b helpers ────────────────────────────
func scanChannelListItems(rows *sql.Rows) ([]store.ChannelListItem, error) {
var result []store.ChannelListItem
for rows.Next() {
var item store.ChannelListItem
var tags, settings string
err := rows.Scan(
&item.ID, &item.UserID, &item.Title, &item.Type, &item.AiMode, &item.Topic,
&item.Description, &item.Model, &item.ProviderConfigID,
&item.SystemPrompt, &item.IsArchived, &item.IsPinned, &item.Folder, &item.FolderID, &item.ProjectID, &item.WorkspaceID,
&tags, &settings,
&item.MessageCount, st(&item.CreatedAtTime), st(&item.UpdatedAtTime),
)
if err != nil {
return nil, err
}
item.Tags = ScanArray(tags)
if item.Tags == nil {
item.Tags = []string{}
}
item.Settings = safeJSONString(settings)
item.CreatedAt = item.CreatedAtTime.Format("2006-01-02T15:04:05Z")
item.UpdatedAt = item.UpdatedAtTime.Format("2006-01-02T15:04:05Z")
result = append(result, item)
}
return result, rows.Err()
}
func safeJSONString(s string) json.RawMessage {
if s == "" || s == "null" {
return json.RawMessage("{}")
}
return json.RawMessage(s)
}

View File

@@ -0,0 +1,143 @@
package sqlite
import (
"context"
"database/sql"
)
// ── Single-field helpers (v0.29.0) ──────────────────────────────────────
func (s *ChannelStore) GetAIMode(ctx context.Context, channelID string) (string, error) {
var aiMode string
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(ai_mode, 'auto') FROM channels WHERE id = ?
`, channelID).Scan(&aiMode)
if err != nil {
return "auto", err
}
return aiMode, nil
}
func (s *ChannelStore) GetTypeAndTeamID(ctx context.Context, channelID string) (string, *string, error) {
var channelType string
var teamID *string
err := DB.QueryRowContext(ctx, `
SELECT COALESCE(type, 'direct'), team_id FROM channels WHERE id = ?
`, channelID).Scan(&channelType, &teamID)
return channelType, teamID, err
}
func (s *ChannelStore) GetSystemPrompt(ctx context.Context, channelID string) (*string, error) {
var prompt *string
err := DB.QueryRowContext(ctx, `
SELECT system_prompt FROM channels WHERE id = ?
`, channelID).Scan(&prompt)
if err == sql.ErrNoRows {
return nil, nil
}
return prompt, err
}
func (s *ChannelStore) GetDefaultModel(ctx context.Context, channelID string) (*string, error) {
var model *string
err := DB.QueryRowContext(ctx, `
SELECT model FROM channels WHERE id = ?
`, channelID).Scan(&model)
if err == sql.ErrNoRows {
return nil, nil
}
return model, err
}
func (s *ChannelStore) TouchUpdatedAt(ctx context.Context, channelID string) error {
_, err := DB.ExecContext(ctx, `UPDATE channels SET updated_at = datetime('now') WHERE id = ?`, channelID)
return err
}
func (s *ChannelStore) ListUserParticipantIDs(ctx context.Context, channelID, excludeUserID string) ([]string, error) {
rows, err := DB.QueryContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = ? AND participant_type = 'user' AND participant_id != ?
`, channelID, excludeUserID)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if ids == nil {
ids = []string{}
}
return ids, rows.Err()
}
func (s *ChannelStore) ListPersonaParticipantIDs(ctx context.Context, channelID string) ([]string, error) {
rows, err := DB.QueryContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = ? AND participant_type = 'persona'
ORDER BY created_at
`, channelID)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if ids == nil {
ids = []string{}
}
return ids, rows.Err()
}
func (s *ChannelStore) GetLeaderPersonaID(ctx context.Context, channelID string) (string, error) {
var leaderID string
err := DB.QueryRowContext(ctx, `
SELECT cp.participant_id
FROM channel_participants cp
JOIN persona_group_members pgm ON pgm.persona_id = cp.participant_id
WHERE cp.channel_id = ?
AND cp.participant_type = 'persona'
AND pgm.is_leader = 1
LIMIT 1
`, channelID).Scan(&leaderID)
if err == nil && leaderID != "" {
return leaderID, nil
}
err = DB.QueryRowContext(ctx, `
SELECT participant_id FROM channel_participants
WHERE channel_id = ? AND participant_type = 'persona'
ORDER BY created_at LIMIT 1
`, channelID).Scan(&leaderID)
if err == sql.ErrNoRows {
return "", nil
}
return leaderID, err
}
func (s *ChannelStore) GetWorkflowInfo(ctx context.Context, channelID string) (*string, int, error) {
var workflowID *string
var currentStage int
err := DB.QueryRowContext(ctx, `
SELECT workflow_id, COALESCE(current_stage, 0)
FROM channels WHERE id = ? AND type = 'workflow'
`, channelID).Scan(&workflowID, &currentStage)
if err == sql.ErrNoRows {
return nil, 0, nil
}
return workflowID, currentStage, err
}

View File

@@ -0,0 +1,164 @@
package sqlite
import (
"context"
"time"
"github.com/google/uuid"
"git.gobha.me/xcaliber/chat-switchboard/database"
"git.gobha.me/xcaliber/chat-switchboard/models"
)
type ExtensionPermissionStore struct{}
func NewExtensionPermissionStore() *ExtensionPermissionStore {
return &ExtensionPermissionStore{}
}
func (s *ExtensionPermissionStore) DeclareForPackage(ctx context.Context, packageID string, permissions []string) error {
tx, err := DB.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if len(permissions) == 0 {
_, err = tx.ExecContext(ctx, `DELETE FROM extension_permissions WHERE package_id = ?`, packageID)
if err != nil {
return err
}
return tx.Commit()
}
declared := make(map[string]bool, len(permissions))
for _, p := range permissions {
declared[p] = true
}
rows, err := tx.QueryContext(ctx,
`SELECT permission FROM extension_permissions WHERE package_id = ?`, packageID)
if err != nil {
return err
}
existing := make(map[string]bool)
for rows.Next() {
var perm string
if err := rows.Scan(&perm); err != nil {
rows.Close()
return err
}
existing[perm] = true
}
rows.Close()
for perm := range existing {
if !declared[perm] {
_, err = tx.ExecContext(ctx,
`DELETE FROM extension_permissions WHERE package_id = ? AND permission = ?`,
packageID, perm)
if err != nil {
return err
}
}
}
for _, perm := range permissions {
if !existing[perm] {
_, err = tx.ExecContext(ctx,
`INSERT OR IGNORE INTO extension_permissions (id, package_id, permission)
VALUES (?, ?, ?)`,
uuid.New().String(), packageID, perm)
if err != nil {
return err
}
}
}
return tx.Commit()
}
func (s *ExtensionPermissionStore) ListForPackage(ctx context.Context, packageID string) ([]models.ExtensionPermission, error) {
rows, err := DB.QueryContext(ctx,
`SELECT id, package_id, permission, granted, granted_by, granted_at, created_at
FROM extension_permissions
WHERE package_id = ?
ORDER BY permission`, packageID)
if err != nil {
return nil, err
}
defer rows.Close()
var perms []models.ExtensionPermission
for rows.Next() {
var p models.ExtensionPermission
if err := rows.Scan(&p.ID, &p.PackageID, &p.Permission, &p.Granted,
&p.GrantedBy, database.SNT(&p.GrantedAt), database.ST(&p.CreatedAt)); err != nil {
return nil, err
}
perms = append(perms, p)
}
if perms == nil {
perms = []models.ExtensionPermission{}
}
return perms, nil
}
func (s *ExtensionPermissionStore) GrantedForPackage(ctx context.Context, packageID string) ([]string, error) {
rows, err := DB.QueryContext(ctx,
`SELECT permission FROM extension_permissions
WHERE package_id = ? AND granted = 1
ORDER BY permission`, packageID)
if err != nil {
return nil, err
}
defer rows.Close()
var perms []string
for rows.Next() {
var p string
if err := rows.Scan(&p); err != nil {
return nil, err
}
perms = append(perms, p)
}
if perms == nil {
perms = []string{}
}
return perms, nil
}
func (s *ExtensionPermissionStore) Grant(ctx context.Context, packageID, permission, grantedBy string) error {
now := time.Now().UTC().Format(time.RFC3339)
_, err := DB.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = 1, granted_by = ?, granted_at = ?
WHERE package_id = ? AND permission = ?`,
grantedBy, now, packageID, permission)
return err
}
func (s *ExtensionPermissionStore) Revoke(ctx context.Context, packageID, permission string) error {
_, err := DB.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = 0, granted_by = NULL, granted_at = NULL
WHERE package_id = ? AND permission = ?`,
packageID, permission)
return err
}
func (s *ExtensionPermissionStore) GrantAll(ctx context.Context, packageID, grantedBy string) error {
now := time.Now().UTC().Format(time.RFC3339)
_, err := DB.ExecContext(ctx,
`UPDATE extension_permissions
SET granted = 1, granted_by = ?, granted_at = ?
WHERE package_id = ? AND granted = 0`,
grantedBy, now, packageID)
return err
}
func (s *ExtensionPermissionStore) DeleteForPackage(ctx context.Context, packageID string) error {
_, err := DB.ExecContext(ctx,
`DELETE FROM extension_permissions WHERE package_id = ?`, packageID)
return err
}

View File

@@ -254,3 +254,12 @@ func (s *FileStore) ListOrphans(ctx context.Context, olderThan time.Duration) ([
} }
return out, rows.Err() return out, rows.Err()
} }
// ── CS5c additions (v0.29.0) ──────────────────────────────────────────
func (s *FileStore) UpdateStorageKey(ctx context.Context, id, key string) error {
_, err := DB.ExecContext(ctx,
`UPDATE files SET storage_key = ?, updated_at = datetime('now') WHERE id = ?`,
key, id)
return err
}

View File

@@ -0,0 +1,82 @@
package sqlite
import (
"context"
"git.gobha.me/xcaliber/chat-switchboard/models"
"git.gobha.me/xcaliber/chat-switchboard/store"
)
type FolderStore struct{}
func NewFolderStore() *FolderStore { return &FolderStore{} }
func (s *FolderStore) List(ctx context.Context, userID string) ([]models.Folder, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, name, parent_id, sort_order, created_at, updated_at
FROM folders WHERE user_id = ?
ORDER BY sort_order, name
`, userID)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.Folder
for rows.Next() {
var f models.Folder
if err := rows.Scan(&f.ID, &f.Name, &f.ParentID, &f.SortOrder,
st(&f.CreatedAt), st(&f.UpdatedAt)); err != nil {
continue
}
f.UserID = userID
result = append(result, f)
}
if result == nil {
result = []models.Folder{}
}
return result, rows.Err()
}
func (s *FolderStore) Create(ctx context.Context, f *models.Folder) error {
f.ID = store.NewID()
_, err := DB.ExecContext(ctx, `
INSERT INTO folders (id, user_id, name, sort_order)
VALUES (?, ?, ?, ?)
`, f.ID, f.UserID, f.Name, f.SortOrder)
if err != nil {
return err
}
return DB.QueryRowContext(ctx, `
SELECT name, parent_id, sort_order, created_at, updated_at
FROM folders WHERE id = ?
`, f.ID).Scan(&f.Name, &f.ParentID, &f.SortOrder, st(&f.CreatedAt), st(&f.UpdatedAt))
}
func (s *FolderStore) Update(ctx context.Context, folderID, userID string, name string, sortOrder *int) (int64, error) {
res, err := DB.ExecContext(ctx, `
UPDATE folders
SET name = COALESCE(NULLIF(?, ''), name),
sort_order = COALESCE(?, sort_order)
WHERE id = ? AND user_id = ?
`, name, sortOrder, folderID, userID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *FolderStore) Delete(ctx context.Context, folderID, userID string) (int64, error) {
res, err := DB.ExecContext(ctx,
`DELETE FROM folders WHERE id = ? AND user_id = ?`, folderID, userID)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func (s *FolderStore) UnassignChannels(ctx context.Context, folderID, userID string) error {
_, err := DB.ExecContext(ctx,
`UPDATE channels SET folder_id = NULL WHERE folder_id = ? AND user_id = ?`, folderID, userID)
return err
}

View File

@@ -56,3 +56,39 @@ func (s *GlobalConfigStore) GetAll(ctx context.Context) (map[string]models.JSONM
} }
return result, rows.Err() return result, rows.Err()
} }
// ── OIDC state (v0.29.0-cs4) ────────────────────────────────────────────
func (s *GlobalConfigStore) SaveOIDCState(ctx context.Context, state, nonce, redirectTo string) error {
_, err := DB.ExecContext(ctx, `
INSERT INTO oidc_auth_state (state, nonce, redirect_to) VALUES (?, ?, ?)
`, state, nonce, redirectTo)
return err
}
func (s *GlobalConfigStore) ConsumeOIDCState(ctx context.Context, state string) (string, string, error) {
var nonce, redirectTo string
err := DB.QueryRowContext(ctx, `
SELECT nonce, COALESCE(redirect_to, '') FROM oidc_auth_state WHERE state = ?
`, state).Scan(&nonce, &redirectTo)
if err != nil {
return "", "", err
}
_, _ = DB.ExecContext(ctx, `DELETE FROM oidc_auth_state WHERE state = ?`, state)
return nonce, redirectTo, nil
}
func (s *GlobalConfigStore) CleanupOIDCState(ctx context.Context) error {
_, err := DB.ExecContext(ctx,
`DELETE FROM oidc_auth_state WHERE created_at < datetime('now', '-10 minutes')`)
return err
}
// ── CS6 additions (v0.29.0) ─────────────────────────────────────────────
func (s *GlobalConfigStore) GetString(ctx context.Context, key string) (string, error) {
var val string
err := DB.QueryRowContext(ctx,
"SELECT value FROM global_settings WHERE key = ?", key).Scan(&val)
return val, err
}

View File

@@ -254,3 +254,34 @@ func (s *MemoryStore) scanMemories(rows *sql.Rows) ([]models.Memory, error) {
// ensure compile-time interface satisfaction // ensure compile-time interface satisfaction
var _ store.MemoryStore = (*MemoryStore)(nil) var _ store.MemoryStore = (*MemoryStore)(nil)
// ── CS5b additions (v0.29.0) ────────────────────────────────────────────
func (s *MemoryStore) SetEmbedding(ctx context.Context, id, embedding string) error {
_, err := DB.ExecContext(ctx,
`UPDATE memories SET embedding = ? WHERE id = ?`, embedding, id)
return err
}
func (s *MemoryStore) GetLastExtractionMessageID(ctx context.Context, channelID, userID string) (string, error) {
var lastID string
err := DB.QueryRowContext(ctx,
`SELECT last_message_id FROM memory_extraction_log WHERE channel_id = ? AND user_id = ?`,
channelID, userID).Scan(&lastID)
if err != nil {
return "", nil
}
return lastID, nil
}
func (s *MemoryStore) UpsertExtractionLog(ctx context.Context, channelID, userID, lastMessageID string, count int) error {
_, err := DB.ExecContext(ctx, `
INSERT INTO memory_extraction_log (id, channel_id, user_id, last_message_id, memory_count)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(channel_id, user_id) DO UPDATE SET
last_message_id = excluded.last_message_id,
extracted_at = datetime('now'),
memory_count = memory_extraction_log.memory_count + excluded.memory_count
`, store.NewID(), channelID, userID, lastMessageID, count)
return err
}

Some files were not shown because too many files have changed in this diff Show More