Changeset 0.9.0 (#50)
This commit is contained in:
82
server/store/postgres/audit.go
Normal file
82
server/store/postgres/audit.go
Normal file
@@ -0,0 +1,82 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type AuditStore struct{}
|
||||
|
||||
func NewAuditStore() *AuditStore { return &AuditStore{} }
|
||||
|
||||
func (s *AuditStore) Log(ctx context.Context, entry *models.AuditEntry) error {
|
||||
metadataJSON := ToJSON(entry.Metadata)
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO audit_log (actor_id, action, resource_type, resource_id, metadata, ip_address, user_agent)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
RETURNING id, created_at`,
|
||||
models.NullString(entry.ActorID), entry.Action, entry.ResourceType,
|
||||
entry.ResourceID, metadataJSON, entry.IPAddress, entry.UserAgent,
|
||||
).Scan(&entry.ID, &entry.CreatedAt)
|
||||
}
|
||||
|
||||
func (s *AuditStore) List(ctx context.Context, opts store.AuditListOptions) ([]models.AuditEntry, int, error) {
|
||||
b := NewSelect("id, actor_id, action, resource_type, resource_id, metadata, ip_address, user_agent, created_at", "audit_log")
|
||||
|
||||
if opts.ActorID != "" {
|
||||
b.Where("actor_id = ?", opts.ActorID)
|
||||
}
|
||||
if opts.Action != "" {
|
||||
b.Where("action = ?", opts.Action)
|
||||
}
|
||||
if opts.ResourceType != "" {
|
||||
b.Where("resource_type = ?", opts.ResourceType)
|
||||
}
|
||||
if opts.ResourceID != "" {
|
||||
b.Where("resource_id = ?", opts.ResourceID)
|
||||
}
|
||||
if opts.Since != nil {
|
||||
b.Where("created_at >= ?", *opts.Since)
|
||||
}
|
||||
if opts.Until != nil {
|
||||
b.Where("created_at <= ?", *opts.Until)
|
||||
}
|
||||
|
||||
// Count
|
||||
countQ, countArgs := b.CountBuild()
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, countQ, countArgs...).Scan(&total)
|
||||
|
||||
// Results
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("created_at", "DESC")
|
||||
}
|
||||
b.Paginate(opts.ListOptions)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.AuditEntry
|
||||
for rows.Next() {
|
||||
var e models.AuditEntry
|
||||
var actorID sql.NullString
|
||||
var metadataJSON []byte
|
||||
err := rows.Scan(&e.ID, &actorID, &e.Action, &e.ResourceType, &e.ResourceID,
|
||||
&metadataJSON, &e.IPAddress, &e.UserAgent, &e.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
e.ActorID = NullableStringPtr(actorID)
|
||||
json.Unmarshal(metadataJSON, &e.Metadata)
|
||||
result = append(result, e)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
224
server/store/postgres/catalog.go
Normal file
224
server/store/postgres/catalog.go
Normal file
@@ -0,0 +1,224 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type CatalogStore struct{}
|
||||
|
||||
func NewCatalogStore() *CatalogStore { return &CatalogStore{} }
|
||||
|
||||
const catalogCols = `id, provider_config_id, model_id, display_name,
|
||||
capabilities, pricing, visibility, last_synced_at, created_at, updated_at`
|
||||
|
||||
// catalogColsMC is catalogCols with mc. prefix for use in JOINs
|
||||
// where id/created_at/updated_at are ambiguous.
|
||||
const catalogColsMC = `mc.id, mc.provider_config_id, mc.model_id, mc.display_name,
|
||||
mc.capabilities, mc.pricing, mc.visibility, mc.last_synced_at, mc.created_at, mc.updated_at`
|
||||
|
||||
// UpsertFromSync bulk-inserts or updates catalog entries from a provider API fetch.
|
||||
// New models default to 'disabled' visibility (secure by default).
|
||||
func (s *CatalogStore) UpsertFromSync(ctx context.Context, providerConfigID string, entries []store.CatalogSyncEntry) (added, updated int, err error) {
|
||||
now := time.Now()
|
||||
for _, e := range entries {
|
||||
capsJSON := ToJSON(e.Capabilities)
|
||||
var pricingJSON []byte
|
||||
if e.Pricing != nil {
|
||||
pricingJSON = ToJSON(e.Pricing)
|
||||
}
|
||||
|
||||
var existingID string
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT id FROM model_catalog WHERE provider_config_id = $1 AND model_id = $2",
|
||||
providerConfigID, e.ModelID,
|
||||
).Scan(&existingID)
|
||||
|
||||
if err == sql.ErrNoRows {
|
||||
// Insert new (disabled by default)
|
||||
_, err = DB.ExecContext(ctx, `
|
||||
INSERT INTO model_catalog (provider_config_id, model_id, display_name,
|
||||
capabilities, pricing, visibility, last_synced_at)
|
||||
VALUES ($1, $2, $3, $4, $5, 'disabled', $6)`,
|
||||
providerConfigID, e.ModelID, e.DisplayName, capsJSON, pricingJSON, now)
|
||||
if err != nil {
|
||||
return added, updated, fmt.Errorf("insert %s: %w", e.ModelID, err)
|
||||
}
|
||||
added++
|
||||
} else if err == nil {
|
||||
// Update existing (preserve visibility)
|
||||
_, err = DB.ExecContext(ctx, `
|
||||
UPDATE model_catalog SET display_name = $1, capabilities = $2,
|
||||
pricing = $3, last_synced_at = $4
|
||||
WHERE id = $5`,
|
||||
e.DisplayName, capsJSON, pricingJSON, now, existingID)
|
||||
if err != nil {
|
||||
return added, updated, fmt.Errorf("update %s: %w", e.ModelID, err)
|
||||
}
|
||||
updated++
|
||||
} else {
|
||||
return added, updated, fmt.Errorf("check %s: %w", e.ModelID, err)
|
||||
}
|
||||
}
|
||||
return added, updated, nil
|
||||
}
|
||||
|
||||
func (s *CatalogStore) GetByID(ctx context.Context, id string) (*models.CatalogEntry, error) {
|
||||
row := DB.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog WHERE id = $1", catalogCols), id)
|
||||
return scanCatalogEntry(row)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) GetByModelID(ctx context.Context, providerConfigID, modelID string) (*models.CatalogEntry, error) {
|
||||
row := DB.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog WHERE provider_config_id = $1 AND model_id = $2", catalogCols),
|
||||
providerConfigID, modelID)
|
||||
return scanCatalogEntry(row)
|
||||
}
|
||||
|
||||
// GetByModelIDAny returns the most recently synced catalog entry for a model_id
|
||||
// across any provider. Used to resolve capabilities for presets with auto-resolve
|
||||
// (no specific provider_config_id).
|
||||
func (s *CatalogStore) GetByModelIDAny(ctx context.Context, modelID string) (*models.CatalogEntry, error) {
|
||||
row := DB.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog WHERE model_id = $1 ORDER BY last_synced_at DESC NULLS LAST LIMIT 1", catalogCols),
|
||||
modelID)
|
||||
return scanCatalogEntry(row)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) ListVisible(ctx context.Context) ([]models.CatalogEntry, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf(`SELECT %s FROM model_catalog mc
|
||||
JOIN provider_configs pc ON pc.id = mc.provider_config_id
|
||||
WHERE mc.visibility = 'enabled' AND pc.scope = 'global' AND pc.is_active = true
|
||||
ORDER BY mc.model_id`, catalogColsMC))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanCatalogEntries(rows)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) ListForProvider(ctx context.Context, providerConfigID string) ([]models.CatalogEntry, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog WHERE provider_config_id = $1 ORDER BY model_id", catalogCols),
|
||||
providerConfigID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanCatalogEntries(rows)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) ListEnabledForProvider(ctx context.Context, providerConfigID string) ([]models.CatalogEntry, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog WHERE provider_config_id = $1 AND visibility = 'enabled' ORDER BY model_id", catalogCols),
|
||||
providerConfigID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanCatalogEntries(rows)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) ListAll(ctx context.Context) ([]models.CatalogEntry, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM model_catalog ORDER BY model_id", catalogCols))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanCatalogEntries(rows)
|
||||
}
|
||||
|
||||
func (s *CatalogStore) SetVisibility(ctx context.Context, id string, visibility string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE model_catalog SET visibility = $1 WHERE id = $2", visibility, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *CatalogStore) BulkSetVisibility(ctx context.Context, providerConfigID string, visibility string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE model_catalog SET visibility = $1 WHERE provider_config_id = $2",
|
||||
visibility, providerConfigID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *CatalogStore) BulkSetVisibilityAll(ctx context.Context, visibility string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE model_catalog SET visibility = $1", visibility)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *CatalogStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM model_catalog WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *CatalogStore) DeleteForProvider(ctx context.Context, providerConfigID string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"DELETE FROM model_catalog WHERE provider_config_id = $1", providerConfigID)
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Scanners ────────────────────────────────
|
||||
|
||||
func scanCatalogEntry(row *sql.Row) (*models.CatalogEntry, error) {
|
||||
var e models.CatalogEntry
|
||||
var capsJSON, pricingJSON []byte
|
||||
var displayName sql.NullString
|
||||
var lastSynced sql.NullTime
|
||||
err := row.Scan(
|
||||
&e.ID, &e.ProviderConfigID, &e.ModelID, &displayName,
|
||||
&capsJSON, &pricingJSON, &e.Visibility, &lastSynced,
|
||||
&e.CreatedAt, &e.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e.DisplayName = NullableString(displayName)
|
||||
json.Unmarshal(capsJSON, &e.Capabilities)
|
||||
if len(pricingJSON) > 0 {
|
||||
e.Pricing = &models.ModelPricing{}
|
||||
json.Unmarshal(pricingJSON, e.Pricing)
|
||||
}
|
||||
if lastSynced.Valid {
|
||||
e.LastSyncedAt = &lastSynced.Time
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
func scanCatalogEntries(rows *sql.Rows) ([]models.CatalogEntry, error) {
|
||||
result := make([]models.CatalogEntry, 0) // never nil — serializes as [] not null
|
||||
for rows.Next() {
|
||||
var e models.CatalogEntry
|
||||
var capsJSON, pricingJSON []byte
|
||||
var displayName sql.NullString
|
||||
var lastSynced sql.NullTime
|
||||
err := rows.Scan(
|
||||
&e.ID, &e.ProviderConfigID, &e.ModelID, &displayName,
|
||||
&capsJSON, &pricingJSON, &e.Visibility, &lastSynced,
|
||||
&e.CreatedAt, &e.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e.DisplayName = NullableString(displayName)
|
||||
json.Unmarshal(capsJSON, &e.Capabilities)
|
||||
if len(pricingJSON) > 0 {
|
||||
e.Pricing = &models.ModelPricing{}
|
||||
json.Unmarshal(pricingJSON, e.Pricing)
|
||||
}
|
||||
if lastSynced.Valid {
|
||||
e.LastSyncedAt = &lastSynced.Time
|
||||
}
|
||||
result = append(result, e)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
221
server/store/postgres/channel.go
Normal file
221
server/store/postgres/channel.go
Normal file
@@ -0,0 +1,221 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type ChannelStore struct{}
|
||||
|
||||
func NewChannelStore() *ChannelStore { return &ChannelStore{} }
|
||||
|
||||
func (s *ChannelStore) Create(ctx context.Context, ch *models.Channel) error {
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO channels (user_id, title, description, type, model, system_prompt,
|
||||
provider_config_id, is_archived, is_pinned, folder_id, team_id, settings)
|
||||
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
ch.UserID, ch.Title, ch.Description, ch.Type, ch.Model, ch.SystemPrompt,
|
||||
models.NullString(ch.ProviderConfigID), ch.IsArchived, ch.IsPinned,
|
||||
models.NullString(ch.FolderID), models.NullString(ch.TeamID), ToJSON(ch.Settings),
|
||||
).Scan(&ch.ID, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *ChannelStore) GetByID(ctx context.Context, id string) (*models.Channel, error) {
|
||||
var ch models.Channel
|
||||
var providerConfigID, folderID, teamID sql.NullString
|
||||
var desc sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, user_id, title, description, type, model, system_prompt,
|
||||
provider_config_id, is_archived, is_pinned, folder_id, team_id, settings,
|
||||
created_at, updated_at
|
||||
FROM channels WHERE id = $1`, id).Scan(
|
||||
&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model, &ch.SystemPrompt,
|
||||
&providerConfigID, &ch.IsArchived, &ch.IsPinned, &folderID, &teamID, &settingsJSON,
|
||||
&ch.CreatedAt, &ch.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ch.Description = NullableString(desc)
|
||||
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
||||
ch.FolderID = NullableStringPtr(folderID)
|
||||
ch.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(settingsJSON, &ch.Settings)
|
||||
return &ch, nil
|
||||
}
|
||||
|
||||
func (s *ChannelStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
b := NewUpdate("channels")
|
||||
for k, v := range fields {
|
||||
if k == "settings" || k == "tags" {
|
||||
b.SetJSON(k, v)
|
||||
} else {
|
||||
b.Set(k, v)
|
||||
}
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ChannelStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM channels WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ChannelStore) ListForUser(ctx context.Context, userID string, opts store.ListOptions) ([]models.Channel, int, error) {
|
||||
// Count
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM channels WHERE user_id = $1", userID).Scan(&total)
|
||||
|
||||
b := NewSelect(
|
||||
"id, user_id, title, description, type, model, system_prompt, provider_config_id, is_archived, is_pinned, folder_id, team_id, settings, created_at, updated_at",
|
||||
"channels",
|
||||
).Where("user_id = ?", userID)
|
||||
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("updated_at", "DESC")
|
||||
}
|
||||
b.Paginate(opts)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Channel
|
||||
for rows.Next() {
|
||||
var ch models.Channel
|
||||
var providerConfigID, folderID, teamID, desc sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model,
|
||||
&ch.SystemPrompt, &providerConfigID, &ch.IsArchived, &ch.IsPinned,
|
||||
&folderID, &teamID, &settingsJSON, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
ch.Description = NullableString(desc)
|
||||
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
||||
ch.FolderID = NullableStringPtr(folderID)
|
||||
ch.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(settingsJSON, &ch.Settings)
|
||||
result = append(result, ch)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
|
||||
func (s *ChannelStore) Search(ctx context.Context, userID, query string, opts store.ListOptions) ([]models.Channel, int, error) {
|
||||
// Simple title search for now — will add full-text when search feature lands
|
||||
b := NewSelect(
|
||||
"id, user_id, title, description, type, model, system_prompt, provider_config_id, is_archived, is_pinned, folder_id, team_id, settings, created_at, updated_at",
|
||||
"channels",
|
||||
).Where("user_id = ?", userID).Where("title ILIKE ?", "%"+query+"%")
|
||||
b.OrderBy("updated_at", "DESC")
|
||||
b.Paginate(opts)
|
||||
|
||||
var total int
|
||||
DB.QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM channels WHERE user_id = $1 AND title ILIKE $2",
|
||||
userID, "%"+query+"%").Scan(&total)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Channel
|
||||
for rows.Next() {
|
||||
var ch models.Channel
|
||||
var providerConfigID, folderID, teamID, desc sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&ch.ID, &ch.UserID, &ch.Title, &desc, &ch.Type, &ch.Model,
|
||||
&ch.SystemPrompt, &providerConfigID, &ch.IsArchived, &ch.IsPinned,
|
||||
&folderID, &teamID, &settingsJSON, &ch.CreatedAt, &ch.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
ch.Description = NullableString(desc)
|
||||
ch.ProviderConfigID = NullableStringPtr(providerConfigID)
|
||||
ch.FolderID = NullableStringPtr(folderID)
|
||||
ch.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(settingsJSON, &ch.Settings)
|
||||
result = append(result, ch)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
|
||||
func (s *ChannelStore) GetCursor(ctx context.Context, channelID, userID string) (*models.ChannelCursor, error) {
|
||||
var c models.ChannelCursor
|
||||
var leafID sql.NullString
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, channel_id, user_id, active_leaf_id, updated_at
|
||||
FROM channel_cursors WHERE channel_id = $1 AND user_id = $2`,
|
||||
channelID, userID).Scan(&c.ID, &c.ChannelID, &c.UserID, &leafID, &c.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.ActiveLeafID = NullableStringPtr(leafID)
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
func (s *ChannelStore) SetCursor(ctx context.Context, channelID, userID, leafID string) error {
|
||||
_, err := 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()`,
|
||||
channelID, userID, leafID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ChannelStore) SetModel(ctx context.Context, cm *models.ChannelModel) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO channel_models (channel_id, model_id, provider_config_id, display_name, system_prompt, settings, is_default)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
ON CONFLICT (channel_id, model_id) DO UPDATE SET
|
||||
provider_config_id = $3, display_name = $4, system_prompt = $5, settings = $6, is_default = $7`,
|
||||
cm.ChannelID, cm.ModelID, cm.ProviderConfigID, cm.DisplayName, cm.SystemPrompt, "{}", cm.IsDefault)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ChannelStore) GetModels(ctx context.Context, channelID string) ([]models.ChannelModel, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
SELECT id, channel_id, model_id, COALESCE(provider_config_id::text, ''),
|
||||
COALESCE(display_name, ''), COALESCE(system_prompt, ''), is_default
|
||||
FROM channel_models WHERE channel_id = $1`, channelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.ChannelModel
|
||||
for rows.Next() {
|
||||
var cm models.ChannelModel
|
||||
if err := rows.Scan(&cm.ID, &cm.ChannelID, &cm.ModelID, &cm.ProviderConfigID,
|
||||
&cm.DisplayName, &cm.SystemPrompt, &cm.IsDefault); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, cm)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *ChannelStore) UserOwns(ctx context.Context, channelID, userID string) (bool, error) {
|
||||
var exists bool
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT EXISTS(SELECT 1 FROM channels WHERE id = $1 AND user_id = $2)",
|
||||
channelID, userID).Scan(&exists)
|
||||
return exists, err
|
||||
}
|
||||
59
server/store/postgres/global_config.go
Normal file
59
server/store/postgres/global_config.go
Normal file
@@ -0,0 +1,59 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type GlobalConfigStore struct{}
|
||||
|
||||
func NewGlobalConfigStore() *GlobalConfigStore { return &GlobalConfigStore{} }
|
||||
|
||||
func (s *GlobalConfigStore) Get(ctx context.Context, key string) (models.JSONMap, error) {
|
||||
var valueJSON []byte
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT value FROM global_settings WHERE key = $1", key).Scan(&valueJSON)
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var result models.JSONMap
|
||||
json.Unmarshal(valueJSON, &result)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *GlobalConfigStore) Set(ctx context.Context, key string, value models.JSONMap, updatedBy string) error {
|
||||
valueJSON := ToJSON(value)
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO global_settings (key, value, updated_by, updated_at)
|
||||
VALUES ($1, $2, $3, NOW())
|
||||
ON CONFLICT (key) DO UPDATE SET value = $2, updated_by = $3, updated_at = NOW()`,
|
||||
key, valueJSON, updatedBy)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *GlobalConfigStore) GetAll(ctx context.Context) (map[string]models.JSONMap, error) {
|
||||
rows, err := DB.QueryContext(ctx, "SELECT key, value FROM global_settings")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
result := make(map[string]models.JSONMap)
|
||||
for rows.Next() {
|
||||
var key string
|
||||
var valueJSON []byte
|
||||
if err := rows.Scan(&key, &valueJSON); err != nil {
|
||||
continue
|
||||
}
|
||||
var m models.JSONMap
|
||||
json.Unmarshal(valueJSON, &m)
|
||||
result[key] = m
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
253
server/store/postgres/helpers.go
Normal file
253
server/store/postgres/helpers.go
Normal file
@@ -0,0 +1,253 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
// DB is the shared database connection pool.
|
||||
// Set during initialization via SetDB.
|
||||
var DB *sql.DB
|
||||
|
||||
// SetDB configures the shared database connection for all stores.
|
||||
func SetDB(db *sql.DB) {
|
||||
DB = db
|
||||
}
|
||||
|
||||
// ── Dynamic SQL Builder ─────────────────────
|
||||
// Replaces the copy-pasted addClause/addField pattern
|
||||
// found in admin.go, presets.go, team_providers.go, apiconfigs.go.
|
||||
|
||||
// UpdateBuilder constructs a dynamic UPDATE statement.
|
||||
type UpdateBuilder struct {
|
||||
table string
|
||||
sets []string
|
||||
args []interface{}
|
||||
where []string
|
||||
argIdx int
|
||||
}
|
||||
|
||||
// NewUpdate creates an UpdateBuilder for the given table.
|
||||
func NewUpdate(table string) *UpdateBuilder {
|
||||
return &UpdateBuilder{table: table}
|
||||
}
|
||||
|
||||
// Set adds a column=value pair to the UPDATE.
|
||||
func (b *UpdateBuilder) Set(col string, val interface{}) *UpdateBuilder {
|
||||
b.argIdx++
|
||||
b.sets = append(b.sets, fmt.Sprintf("%s = $%d", col, b.argIdx))
|
||||
b.args = append(b.args, val)
|
||||
return b
|
||||
}
|
||||
|
||||
// SetJSON adds a JSONB column from a map.
|
||||
func (b *UpdateBuilder) SetJSON(col string, val interface{}) *UpdateBuilder {
|
||||
data, err := json.Marshal(val)
|
||||
if err != nil {
|
||||
data = []byte("{}")
|
||||
}
|
||||
return b.Set(col, string(data))
|
||||
}
|
||||
|
||||
// SetIf conditionally adds a column if the pointer is non-nil.
|
||||
func (b *UpdateBuilder) SetIf(col string, val interface{}, set bool) *UpdateBuilder {
|
||||
if set {
|
||||
return b.Set(col, val)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// Where adds a WHERE condition.
|
||||
func (b *UpdateBuilder) Where(col string, val interface{}) *UpdateBuilder {
|
||||
b.argIdx++
|
||||
b.where = append(b.where, fmt.Sprintf("%s = $%d", col, b.argIdx))
|
||||
b.args = append(b.args, val)
|
||||
return b
|
||||
}
|
||||
|
||||
// HasSets returns true if any SET clauses were added.
|
||||
func (b *UpdateBuilder) HasSets() bool {
|
||||
return len(b.sets) > 0
|
||||
}
|
||||
|
||||
// Build returns the SQL string and args.
|
||||
func (b *UpdateBuilder) Build() (string, []interface{}) {
|
||||
sql := fmt.Sprintf("UPDATE %s SET %s", b.table, strings.Join(b.sets, ", "))
|
||||
if len(b.where) > 0 {
|
||||
sql += " WHERE " + strings.Join(b.where, " AND ")
|
||||
}
|
||||
return sql, b.args
|
||||
}
|
||||
|
||||
// Exec executes the built UPDATE.
|
||||
func (b *UpdateBuilder) Exec(db *sql.DB) (sql.Result, error) {
|
||||
q, args := b.Build()
|
||||
return db.Exec(q, args...)
|
||||
}
|
||||
|
||||
// ── Query Builder ───────────────────────────
|
||||
|
||||
// SelectBuilder constructs a dynamic SELECT statement.
|
||||
type SelectBuilder struct {
|
||||
cols string
|
||||
table string
|
||||
joins []string
|
||||
where []string
|
||||
args []interface{}
|
||||
orderBy string
|
||||
limit int
|
||||
offset int
|
||||
argIdx int
|
||||
}
|
||||
|
||||
// NewSelect creates a SelectBuilder.
|
||||
func NewSelect(cols, table string) *SelectBuilder {
|
||||
return &SelectBuilder{cols: cols, table: table}
|
||||
}
|
||||
|
||||
// Join adds a JOIN clause.
|
||||
func (b *SelectBuilder) Join(join string) *SelectBuilder {
|
||||
b.joins = append(b.joins, join)
|
||||
return b
|
||||
}
|
||||
|
||||
// Where adds a WHERE condition with a parameter.
|
||||
func (b *SelectBuilder) Where(clause string, args ...interface{}) *SelectBuilder {
|
||||
for _, arg := range args {
|
||||
b.argIdx++
|
||||
clause = strings.Replace(clause, "?", fmt.Sprintf("$%d", b.argIdx), 1)
|
||||
b.args = append(b.args, arg)
|
||||
}
|
||||
b.where = append(b.where, clause)
|
||||
return b
|
||||
}
|
||||
|
||||
// WhereRaw adds a WHERE condition without parameters.
|
||||
func (b *SelectBuilder) WhereRaw(clause string) *SelectBuilder {
|
||||
b.where = append(b.where, clause)
|
||||
return b
|
||||
}
|
||||
|
||||
// OrderBy sets the ORDER BY clause.
|
||||
func (b *SelectBuilder) OrderBy(col, order string) *SelectBuilder {
|
||||
if order == "" {
|
||||
order = "DESC"
|
||||
}
|
||||
b.orderBy = fmt.Sprintf("%s %s", col, strings.ToUpper(order))
|
||||
return b
|
||||
}
|
||||
|
||||
// Paginate sets LIMIT and OFFSET from ListOptions.
|
||||
func (b *SelectBuilder) Paginate(opts store.ListOptions) *SelectBuilder {
|
||||
if opts.Limit > 0 {
|
||||
b.limit = opts.Limit
|
||||
}
|
||||
if opts.Offset > 0 {
|
||||
b.offset = opts.Offset
|
||||
}
|
||||
if opts.Sort != "" {
|
||||
b.OrderBy(opts.Sort, opts.Order)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// Build returns the SQL string and args.
|
||||
func (b *SelectBuilder) Build() (string, []interface{}) {
|
||||
q := fmt.Sprintf("SELECT %s FROM %s", b.cols, b.table)
|
||||
for _, j := range b.joins {
|
||||
q += " " + j
|
||||
}
|
||||
if len(b.where) > 0 {
|
||||
q += " WHERE " + strings.Join(b.where, " AND ")
|
||||
}
|
||||
if b.orderBy != "" {
|
||||
q += " ORDER BY " + b.orderBy
|
||||
}
|
||||
if b.limit > 0 {
|
||||
q += fmt.Sprintf(" LIMIT %d", b.limit)
|
||||
}
|
||||
if b.offset > 0 {
|
||||
q += fmt.Sprintf(" OFFSET %d", b.offset)
|
||||
}
|
||||
return q, b.args
|
||||
}
|
||||
|
||||
// CountBuild returns a SELECT COUNT(*) version of the query (no order/limit).
|
||||
func (b *SelectBuilder) CountBuild() (string, []interface{}) {
|
||||
q := fmt.Sprintf("SELECT COUNT(*) FROM %s", b.table)
|
||||
for _, j := range b.joins {
|
||||
q += " " + j
|
||||
}
|
||||
if len(b.where) > 0 {
|
||||
q += " WHERE " + strings.Join(b.where, " AND ")
|
||||
}
|
||||
return q, b.args
|
||||
}
|
||||
|
||||
// ── JSONB Helpers ───────────────────────────
|
||||
|
||||
// ToJSON marshals a value to JSON bytes for JSONB columns.
|
||||
func ToJSON(v interface{}) []byte {
|
||||
if v == nil {
|
||||
return []byte("{}")
|
||||
}
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return []byte("{}")
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
// ScanJSON scans a JSONB column into a target.
|
||||
func ScanJSON(src interface{}, dst interface{}) error {
|
||||
if src == nil {
|
||||
return nil
|
||||
}
|
||||
var data []byte
|
||||
switch v := src.(type) {
|
||||
case []byte:
|
||||
data = v
|
||||
case string:
|
||||
data = []byte(v)
|
||||
default:
|
||||
return fmt.Errorf("unsupported JSONB type: %T", src)
|
||||
}
|
||||
return json.Unmarshal(data, dst)
|
||||
}
|
||||
|
||||
// NullableString returns the string value or empty string from sql.NullString.
|
||||
func NullableString(ns sql.NullString) string {
|
||||
if ns.Valid {
|
||||
return ns.String
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// NullableStringPtr returns a *string from sql.NullString (nil if not valid).
|
||||
func NullableStringPtr(ns sql.NullString) *string {
|
||||
if ns.Valid {
|
||||
return &ns.String
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// NullableFloat64Ptr returns a *float64 from sql.NullFloat64.
|
||||
func NullableFloat64Ptr(nf sql.NullFloat64) *float64 {
|
||||
if nf.Valid {
|
||||
return &nf.Float64
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// NullableIntPtr returns an *int from sql.NullInt64.
|
||||
func NullableIntPtr(ni sql.NullInt64) *int {
|
||||
if ni.Valid {
|
||||
v := int(ni.Int64)
|
||||
return &v
|
||||
}
|
||||
return nil
|
||||
}
|
||||
191
server/store/postgres/message.go
Normal file
191
server/store/postgres/message.go
Normal file
@@ -0,0 +1,191 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type MessageStore struct{}
|
||||
|
||||
func NewMessageStore() *MessageStore { return &MessageStore{} }
|
||||
|
||||
func (s *MessageStore) Create(ctx context.Context, m *models.Message) error {
|
||||
toolCallsJSON := ToJSON(m.ToolCalls)
|
||||
metadataJSON := ToJSON(m.Metadata)
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO messages (channel_id, role, content, model, tokens_used, tool_calls,
|
||||
metadata, parent_id, sibling_index, participant_type, participant_id)
|
||||
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
|
||||
RETURNING id, created_at`,
|
||||
m.ChannelID, m.Role, m.Content, m.Model, m.TokensUsed,
|
||||
toolCallsJSON, metadataJSON,
|
||||
models.NullString(m.ParentID), m.SiblingIndex,
|
||||
m.ParticipantType, m.ParticipantID,
|
||||
).Scan(&m.ID, &m.CreatedAt)
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetByID(ctx context.Context, id string) (*models.Message, error) {
|
||||
var m models.Message
|
||||
var parentID sql.NullString
|
||||
var toolCallsJSON, metadataJSON []byte
|
||||
var deletedAt sql.NullTime
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
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 messages WHERE id = $1`, id).Scan(
|
||||
&m.ID, &m.ChannelID, &m.Role, &m.Content, &m.Model, &m.TokensUsed,
|
||||
&toolCallsJSON, &metadataJSON,
|
||||
&parentID, &m.SiblingIndex, &m.ParticipantType, &m.ParticipantID,
|
||||
&deletedAt, &m.CreatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m.ParentID = NullableStringPtr(parentID)
|
||||
json.Unmarshal(toolCallsJSON, &m.ToolCalls)
|
||||
json.Unmarshal(metadataJSON, &m.Metadata)
|
||||
if deletedAt.Valid {
|
||||
m.DeletedAt = &deletedAt.Time
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
b := NewUpdate("messages")
|
||||
for k, v := range fields {
|
||||
if k == "tool_calls" || k == "metadata" {
|
||||
b.SetJSON(k, v)
|
||||
} else {
|
||||
b.Set(k, v)
|
||||
}
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *MessageStore) Delete(ctx context.Context, id string) error {
|
||||
now := time.Now()
|
||||
_, err := DB.ExecContext(ctx, "UPDATE messages SET deleted_at = $1 WHERE id = $2", now, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *MessageStore) ListForChannel(ctx context.Context, channelID string, opts store.ListOptions) ([]models.Message, error) {
|
||||
b := NewSelect(
|
||||
"id, channel_id, role, content, model, tokens_used, tool_calls, metadata, parent_id, sibling_index, participant_type, participant_id, deleted_at, created_at",
|
||||
"messages",
|
||||
).Where("channel_id = ?", channelID).WhereRaw("deleted_at IS NULL")
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("created_at", "ASC")
|
||||
}
|
||||
b.Paginate(opts)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetChildren(ctx context.Context, parentID string) ([]models.Message, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
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 messages WHERE parent_id = $1 AND deleted_at IS NULL
|
||||
ORDER BY sibling_index`, parentID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetSiblings(ctx context.Context, messageID string) ([]models.Message, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
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 messages
|
||||
WHERE parent_id = (SELECT parent_id FROM messages WHERE id = $1)
|
||||
AND deleted_at IS NULL
|
||||
ORDER BY sibling_index`, messageID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetPathToRoot(ctx context.Context, messageID string) ([]models.Message, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
WITH RECURSIVE path AS (
|
||||
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 messages WHERE id = $1
|
||||
UNION ALL
|
||||
SELECT m.id, m.channel_id, m.role, m.content, m.model, m.tokens_used, m.tool_calls, m.metadata,
|
||||
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
|
||||
)
|
||||
SELECT * FROM path ORDER BY created_at ASC`, messageID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanMessages(rows)
|
||||
}
|
||||
|
||||
func (s *MessageStore) GetNextSiblingIndex(ctx context.Context, parentID string) (int, error) {
|
||||
var maxIdx sql.NullInt64
|
||||
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, err
|
||||
}
|
||||
return int(maxIdx.Int64) + 1, nil
|
||||
}
|
||||
|
||||
func (s *MessageStore) CountForChannel(ctx context.Context, channelID string) (int, error) {
|
||||
var count int
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM messages WHERE channel_id = $1 AND deleted_at IS NULL",
|
||||
channelID).Scan(&count)
|
||||
return count, err
|
||||
}
|
||||
|
||||
func scanMessages(rows *sql.Rows) ([]models.Message, error) {
|
||||
var result []models.Message
|
||||
for rows.Next() {
|
||||
var m models.Message
|
||||
var parentID sql.NullString
|
||||
var toolCallsJSON, metadataJSON []byte
|
||||
var deletedAt sql.NullTime
|
||||
err := rows.Scan(
|
||||
&m.ID, &m.ChannelID, &m.Role, &m.Content, &m.Model, &m.TokensUsed,
|
||||
&toolCallsJSON, &metadataJSON,
|
||||
&parentID, &m.SiblingIndex, &m.ParticipantType, &m.ParticipantID,
|
||||
&deletedAt, &m.CreatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m.ParentID = NullableStringPtr(parentID)
|
||||
json.Unmarshal(toolCallsJSON, &m.ToolCalls)
|
||||
json.Unmarshal(metadataJSON, &m.Metadata)
|
||||
if deletedAt.Valid {
|
||||
m.DeletedAt = &deletedAt.Time
|
||||
}
|
||||
result = append(result, m)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
191
server/store/postgres/note.go
Normal file
191
server/store/postgres/note.go
Normal file
@@ -0,0 +1,191 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
|
||||
"github.com/lib/pq"
|
||||
)
|
||||
|
||||
type NoteStore struct{}
|
||||
|
||||
func NewNoteStore() *NoteStore { return &NoteStore{} }
|
||||
|
||||
func (s *NoteStore) Create(ctx context.Context, n *models.Note) error {
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO notes (user_id, title, content, folder_path, tags, metadata, source_channel_id, team_id)
|
||||
VALUES ($1,$2,$3,$4,$5,$6,$7,$8)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
n.UserID, n.Title, n.Content, n.FolderPath,
|
||||
pq.Array(n.Tags), ToJSON(n.Metadata),
|
||||
models.NullString(n.SourceChannelID), models.NullString(n.TeamID),
|
||||
).Scan(&n.ID, &n.CreatedAt, &n.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *NoteStore) GetByID(ctx context.Context, id string) (*models.Note, error) {
|
||||
var n models.Note
|
||||
var sourceChannelID, teamID sql.NullString
|
||||
var metadataJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, user_id, title, content, folder_path, tags, metadata,
|
||||
source_channel_id, team_id, created_at, updated_at
|
||||
FROM notes WHERE id = $1`, id).Scan(
|
||||
&n.ID, &n.UserID, &n.Title, &n.Content, &n.FolderPath,
|
||||
pq.Array(&n.Tags), &metadataJSON,
|
||||
&sourceChannelID, &teamID, &n.CreatedAt, &n.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
n.SourceChannelID = NullableStringPtr(sourceChannelID)
|
||||
n.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(metadataJSON, &n.Metadata)
|
||||
return &n, nil
|
||||
}
|
||||
|
||||
func (s *NoteStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
b := NewUpdate("notes")
|
||||
for k, v := range fields {
|
||||
if k == "metadata" {
|
||||
b.SetJSON(k, v)
|
||||
} else if k == "tags" {
|
||||
if tags, ok := v.([]string); ok {
|
||||
b.Set(k, pq.Array(tags))
|
||||
}
|
||||
} else {
|
||||
b.Set(k, v)
|
||||
}
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *NoteStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM notes WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *NoteStore) ListForUser(ctx context.Context, userID string, opts store.NoteListOptions) ([]models.Note, int, error) {
|
||||
b := NewSelect(
|
||||
"id, user_id, title, content, folder_path, tags, metadata, source_channel_id, team_id, created_at, updated_at",
|
||||
"notes",
|
||||
).Where("user_id = ?", userID)
|
||||
|
||||
if opts.FolderPath != "" {
|
||||
b.Where("folder_path = ?", opts.FolderPath)
|
||||
}
|
||||
if opts.Tag != "" {
|
||||
b.Where("? = ANY(tags)", opts.Tag)
|
||||
}
|
||||
if opts.TeamID != "" {
|
||||
b.Where("team_id = ?", opts.TeamID)
|
||||
}
|
||||
|
||||
// Count
|
||||
countQ, countArgs := b.CountBuild()
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, countQ, countArgs...).Scan(&total)
|
||||
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("updated_at", "DESC")
|
||||
}
|
||||
b.Paginate(opts.ListOptions)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Note
|
||||
for rows.Next() {
|
||||
var n models.Note
|
||||
var sourceChannelID, teamID sql.NullString
|
||||
var metadataJSON []byte
|
||||
err := rows.Scan(&n.ID, &n.UserID, &n.Title, &n.Content, &n.FolderPath,
|
||||
pq.Array(&n.Tags), &metadataJSON,
|
||||
&sourceChannelID, &teamID, &n.CreatedAt, &n.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
n.SourceChannelID = NullableStringPtr(sourceChannelID)
|
||||
n.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(metadataJSON, &n.Metadata)
|
||||
result = append(result, n)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
|
||||
func (s *NoteStore) Search(ctx context.Context, userID, query string, opts store.ListOptions) ([]models.Note, int, error) {
|
||||
tsQuery := strings.Join(strings.Fields(query), " & ")
|
||||
|
||||
b := NewSelect(
|
||||
"id, user_id, title, content, folder_path, tags, metadata, source_channel_id, team_id, created_at, updated_at",
|
||||
"notes",
|
||||
).Where("user_id = ?", userID).Where("search_vector @@ to_tsquery('english', ?)", tsQuery)
|
||||
b.OrderBy("ts_rank(search_vector, to_tsquery('english', '"+tsQuery+"'))", "DESC")
|
||||
b.Paginate(opts)
|
||||
|
||||
var total int
|
||||
DB.QueryRowContext(ctx,
|
||||
"SELECT COUNT(*) FROM notes WHERE user_id = $1 AND search_vector @@ to_tsquery('english', $2)",
|
||||
userID, tsQuery).Scan(&total)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Note
|
||||
for rows.Next() {
|
||||
var n models.Note
|
||||
var sourceChannelID, teamID sql.NullString
|
||||
var metadataJSON []byte
|
||||
err := rows.Scan(&n.ID, &n.UserID, &n.Title, &n.Content, &n.FolderPath,
|
||||
pq.Array(&n.Tags), &metadataJSON,
|
||||
&sourceChannelID, &teamID, &n.CreatedAt, &n.UpdatedAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
n.SourceChannelID = NullableStringPtr(sourceChannelID)
|
||||
n.TeamID = NullableStringPtr(teamID)
|
||||
json.Unmarshal(metadataJSON, &n.Metadata)
|
||||
result = append(result, n)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
|
||||
func (s *NoteStore) BulkDelete(ctx context.Context, ids []string, userID string) (int, error) {
|
||||
if len(ids) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
placeholders := make([]string, len(ids))
|
||||
args := make([]interface{}, 0, len(ids)+1)
|
||||
args = append(args, userID)
|
||||
for i, id := range ids {
|
||||
placeholders[i] = fmt.Sprintf("$%d", i+2)
|
||||
args = append(args, id)
|
||||
}
|
||||
result, err := DB.ExecContext(ctx,
|
||||
fmt.Sprintf("DELETE FROM notes WHERE user_id = $1 AND id IN (%s)",
|
||||
strings.Join(placeholders, ",")),
|
||||
args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
n, _ := result.RowsAffected()
|
||||
return int(n), nil
|
||||
}
|
||||
295
server/store/postgres/persona.go
Normal file
295
server/store/postgres/persona.go
Normal file
@@ -0,0 +1,295 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type PersonaStore struct{}
|
||||
|
||||
func NewPersonaStore() *PersonaStore { return &PersonaStore{} }
|
||||
|
||||
const personaCols = `id, name, description, icon, avatar, base_model_id, provider_config_id,
|
||||
system_prompt, temperature, max_tokens, thinking_budget, top_p,
|
||||
scope, owner_id, created_by, is_active, is_shared, created_at, updated_at`
|
||||
|
||||
func (s *PersonaStore) Create(ctx context.Context, p *models.Persona) error {
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO personas (name, description, icon, avatar, base_model_id, provider_config_id,
|
||||
system_prompt, temperature, max_tokens, thinking_budget, top_p,
|
||||
scope, owner_id, created_by, is_active, is_shared)
|
||||
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
p.Name, p.Description, p.Icon, p.Avatar, p.BaseModelID,
|
||||
models.NullString(p.ProviderConfigID),
|
||||
p.SystemPrompt, models.NullFloat(p.Temperature), models.NullInt(p.MaxTokens),
|
||||
models.NullInt(p.ThinkingBudget), models.NullFloat(p.TopP),
|
||||
p.Scope, models.NullString(p.OwnerID), p.CreatedBy, p.IsActive, p.IsShared,
|
||||
).Scan(&p.ID, &p.CreatedAt, &p.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *PersonaStore) GetByID(ctx context.Context, id string) (*models.Persona, error) {
|
||||
row := DB.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM personas WHERE id = $1", personaCols), id)
|
||||
p, err := scanPersona(row)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Load grants
|
||||
grants, _ := s.GetGrants(ctx, id)
|
||||
p.Grants = grants
|
||||
return p, nil
|
||||
}
|
||||
|
||||
func (s *PersonaStore) Update(ctx context.Context, id string, patch models.PersonaPatch) error {
|
||||
b := NewUpdate("personas")
|
||||
if patch.Name != nil {
|
||||
b.Set("name", *patch.Name)
|
||||
}
|
||||
if patch.Description != nil {
|
||||
b.Set("description", *patch.Description)
|
||||
}
|
||||
if patch.Icon != nil {
|
||||
b.Set("icon", *patch.Icon)
|
||||
}
|
||||
if patch.Avatar != nil {
|
||||
b.Set("avatar", *patch.Avatar)
|
||||
}
|
||||
if patch.BaseModelID != nil {
|
||||
b.Set("base_model_id", *patch.BaseModelID)
|
||||
}
|
||||
if patch.ProviderConfigID != nil {
|
||||
b.Set("provider_config_id", models.NullString(patch.ProviderConfigID))
|
||||
}
|
||||
if patch.SystemPrompt != nil {
|
||||
b.Set("system_prompt", *patch.SystemPrompt)
|
||||
}
|
||||
if patch.Temperature != nil {
|
||||
b.Set("temperature", models.NullFloat(patch.Temperature))
|
||||
}
|
||||
if patch.MaxTokens != nil {
|
||||
b.Set("max_tokens", models.NullInt(patch.MaxTokens))
|
||||
}
|
||||
if patch.ThinkingBudget != nil {
|
||||
b.Set("thinking_budget", models.NullInt(patch.ThinkingBudget))
|
||||
}
|
||||
if patch.TopP != nil {
|
||||
b.Set("top_p", models.NullFloat(patch.TopP))
|
||||
}
|
||||
if patch.IsActive != nil {
|
||||
b.Set("is_active", *patch.IsActive)
|
||||
}
|
||||
if patch.IsShared != nil {
|
||||
b.Set("is_shared", *patch.IsShared)
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *PersonaStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM personas WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
// ListForUser returns all Personas visible to a user:
|
||||
// global active + team-scoped (for user's teams) + personal + shared.
|
||||
func (s *PersonaStore) ListForUser(ctx context.Context, userID string) ([]models.Persona, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf(`SELECT %s FROM personas WHERE is_active = true AND (
|
||||
scope = 'global'
|
||||
OR (scope = 'personal' AND created_by = $1)
|
||||
OR (scope = 'team' AND owner_id IN (
|
||||
SELECT team_id FROM team_members WHERE user_id = $1
|
||||
))
|
||||
OR (scope = 'personal' AND is_shared = true)
|
||||
) ORDER BY scope, name`, personaCols), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanPersonas(rows)
|
||||
}
|
||||
|
||||
func (s *PersonaStore) ListForTeam(ctx context.Context, teamID string) ([]models.Persona, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM personas WHERE scope = 'team' AND owner_id = $1 AND is_active = true ORDER BY name", personaCols),
|
||||
teamID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanPersonas(rows)
|
||||
}
|
||||
|
||||
func (s *PersonaStore) ListGlobal(ctx context.Context) ([]models.Persona, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM personas WHERE scope = 'global' ORDER BY name", personaCols))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanPersonas(rows)
|
||||
}
|
||||
|
||||
func (s *PersonaStore) ListPersonal(ctx context.Context, userID string) ([]models.Persona, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM personas WHERE scope = 'personal' AND created_by = $1 ORDER BY name", personaCols),
|
||||
userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanPersonas(rows)
|
||||
}
|
||||
|
||||
// ── Grants ──────────────────────────────────
|
||||
|
||||
// SetGrants replaces all grants for a Persona (delete + re-insert).
|
||||
func (s *PersonaStore) SetGrants(ctx context.Context, personaID string, grants []models.Grant) error {
|
||||
tx, err := DB.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
|
||||
_, err = tx.ExecContext(ctx, "DELETE FROM persona_grants WHERE persona_id = $1", personaID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, g := range grants {
|
||||
configJSON := ToJSON(g.Config)
|
||||
_, err = tx.ExecContext(ctx, `
|
||||
INSERT INTO persona_grants (persona_id, grant_type, grant_ref, config)
|
||||
VALUES ($1, $2, $3, $4)`,
|
||||
personaID, g.GrantType, g.GrantRef, configJSON)
|
||||
if err != nil {
|
||||
return fmt.Errorf("grant %s/%s: %w", g.GrantType, g.GrantRef, err)
|
||||
}
|
||||
}
|
||||
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (s *PersonaStore) GetGrants(ctx context.Context, personaID string) ([]models.Grant, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
"SELECT id, persona_id, grant_type, grant_ref, config, created_at FROM persona_grants WHERE persona_id = $1 ORDER BY grant_type, grant_ref",
|
||||
personaID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Grant
|
||||
for rows.Next() {
|
||||
var g models.Grant
|
||||
var configJSON []byte
|
||||
err := rows.Scan(&g.ID, &g.PersonaID, &g.GrantType, &g.GrantRef, &configJSON, &g.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
json.Unmarshal(configJSON, &g.Config)
|
||||
result = append(result, g)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
// GetToolGrants returns just the tool names for a Persona.
|
||||
func (s *PersonaStore) GetToolGrants(ctx context.Context, personaID string) ([]string, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
"SELECT grant_ref FROM persona_grants WHERE persona_id = $1 AND grant_type = 'tool' ORDER BY grant_ref",
|
||||
personaID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []string
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, name)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
// UserCanAccess checks if a user can see/use a specific Persona.
|
||||
func (s *PersonaStore) UserCanAccess(ctx context.Context, userID, personaID string) (bool, error) {
|
||||
var exists bool
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT EXISTS(
|
||||
SELECT 1 FROM personas WHERE id = $2 AND is_active = true AND (
|
||||
scope = 'global'
|
||||
OR (scope = 'personal' AND created_by = $1)
|
||||
OR (scope = 'team' AND owner_id IN (
|
||||
SELECT team_id FROM team_members WHERE user_id = $1
|
||||
))
|
||||
OR (scope = 'personal' AND is_shared = true)
|
||||
)
|
||||
)`, userID, personaID).Scan(&exists)
|
||||
return exists, err
|
||||
}
|
||||
|
||||
// ── Scanners ────────────────────────────────
|
||||
|
||||
func scanPersona(row *sql.Row) (*models.Persona, error) {
|
||||
var p models.Persona
|
||||
var providerConfigID, ownerID sql.NullString
|
||||
var temp, topP sql.NullFloat64
|
||||
var maxTokens, thinkingBudget sql.NullInt64
|
||||
err := row.Scan(
|
||||
&p.ID, &p.Name, &p.Description, &p.Icon, &p.Avatar,
|
||||
&p.BaseModelID, &providerConfigID,
|
||||
&p.SystemPrompt, &temp, &maxTokens, &thinkingBudget, &topP,
|
||||
&p.Scope, &ownerID, &p.CreatedBy, &p.IsActive, &p.IsShared,
|
||||
&p.CreatedAt, &p.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.ProviderConfigID = NullableStringPtr(providerConfigID)
|
||||
p.OwnerID = NullableStringPtr(ownerID)
|
||||
p.Temperature = NullableFloat64Ptr(temp)
|
||||
p.MaxTokens = NullableIntPtr(maxTokens)
|
||||
p.ThinkingBudget = NullableIntPtr(thinkingBudget)
|
||||
p.TopP = NullableFloat64Ptr(topP)
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func scanPersonas(rows *sql.Rows) ([]models.Persona, error) {
|
||||
var result []models.Persona
|
||||
for rows.Next() {
|
||||
var p models.Persona
|
||||
var providerConfigID, ownerID sql.NullString
|
||||
var temp, topP sql.NullFloat64
|
||||
var maxTokens, thinkingBudget sql.NullInt64
|
||||
err := rows.Scan(
|
||||
&p.ID, &p.Name, &p.Description, &p.Icon, &p.Avatar,
|
||||
&p.BaseModelID, &providerConfigID,
|
||||
&p.SystemPrompt, &temp, &maxTokens, &thinkingBudget, &topP,
|
||||
&p.Scope, &ownerID, &p.CreatedBy, &p.IsActive, &p.IsShared,
|
||||
&p.CreatedAt, &p.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.ProviderConfigID = NullableStringPtr(providerConfigID)
|
||||
p.OwnerID = NullableStringPtr(ownerID)
|
||||
p.Temperature = NullableFloat64Ptr(temp)
|
||||
p.MaxTokens = NullableIntPtr(maxTokens)
|
||||
p.ThinkingBudget = NullableIntPtr(thinkingBudget)
|
||||
p.TopP = NullableFloat64Ptr(topP)
|
||||
result = append(result, p)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
68
server/store/postgres/policy.go
Normal file
68
server/store/postgres/policy.go
Normal file
@@ -0,0 +1,68 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type PolicyStore struct{}
|
||||
|
||||
func NewPolicyStore() *PolicyStore { return &PolicyStore{} }
|
||||
|
||||
// Get returns a policy value by key. Falls back to PolicyDefaults if not in DB.
|
||||
func (s *PolicyStore) Get(ctx context.Context, key string) (string, error) {
|
||||
var value string
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT value FROM platform_policies WHERE key = $1", key).Scan(&value)
|
||||
if err == sql.ErrNoRows {
|
||||
if def, ok := models.PolicyDefaults[key]; ok {
|
||||
return def, nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
return value, err
|
||||
}
|
||||
|
||||
// GetBool returns a policy value as a boolean.
|
||||
func (s *PolicyStore) GetBool(ctx context.Context, key string) (bool, error) {
|
||||
val, err := s.Get(ctx, key)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return val == "true", nil
|
||||
}
|
||||
|
||||
// Set upserts a policy value.
|
||||
func (s *PolicyStore) Set(ctx context.Context, key, value, updatedBy string) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO platform_policies (key, value, updated_by, updated_at)
|
||||
VALUES ($1, $2, $3, NOW())
|
||||
ON CONFLICT (key) DO UPDATE SET value = $2, updated_by = $3, updated_at = NOW()`,
|
||||
key, value, updatedBy)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetAll returns all platform policies, merged with defaults.
|
||||
func (s *PolicyStore) GetAll(ctx context.Context) (map[string]string, error) {
|
||||
result := make(map[string]string)
|
||||
// Start with defaults
|
||||
for k, v := range models.PolicyDefaults {
|
||||
result[k] = v
|
||||
}
|
||||
// Override with DB values
|
||||
rows, err := DB.QueryContext(ctx, "SELECT key, value FROM platform_policies")
|
||||
if err != nil {
|
||||
return result, err // return defaults on error
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var k, v string
|
||||
if err := rows.Scan(&k, &v); err != nil {
|
||||
continue
|
||||
}
|
||||
result[k] = v
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
194
server/store/postgres/provider.go
Normal file
194
server/store/postgres/provider.go
Normal file
@@ -0,0 +1,194 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type ProviderStore struct{}
|
||||
|
||||
func NewProviderStore() *ProviderStore { return &ProviderStore{} }
|
||||
|
||||
const providerCols = `id, scope, owner_id, name, provider, endpoint, api_key_enc,
|
||||
model_default, config, headers, settings, is_active, is_private, created_at, updated_at`
|
||||
|
||||
func (s *ProviderStore) Create(ctx context.Context, cfg *models.ProviderConfig) error {
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO provider_configs (scope, owner_id, name, provider, endpoint, api_key_enc,
|
||||
model_default, config, headers, settings, is_active, is_private)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
cfg.Scope, models.NullString(cfg.OwnerID), cfg.Name, cfg.Provider, cfg.Endpoint,
|
||||
cfg.APIKeyEnc, cfg.ModelDefault,
|
||||
ToJSON(cfg.Config), ToJSON(cfg.Headers), ToJSON(cfg.Settings),
|
||||
cfg.IsActive, cfg.IsPrivate,
|
||||
).Scan(&cfg.ID, &cfg.CreatedAt, &cfg.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *ProviderStore) GetByID(ctx context.Context, id string) (*models.ProviderConfig, error) {
|
||||
row := DB.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM provider_configs WHERE id = $1", providerCols), id)
|
||||
return scanProvider(row)
|
||||
}
|
||||
|
||||
func (s *ProviderStore) Update(ctx context.Context, id string, patch models.ProviderConfigPatch) error {
|
||||
b := NewUpdate("provider_configs")
|
||||
if patch.Name != nil {
|
||||
b.Set("name", *patch.Name)
|
||||
}
|
||||
if patch.Endpoint != nil {
|
||||
b.Set("endpoint", *patch.Endpoint)
|
||||
}
|
||||
if patch.APIKeyEnc != nil {
|
||||
b.Set("api_key_enc", *patch.APIKeyEnc)
|
||||
}
|
||||
if patch.ModelDefault != nil {
|
||||
b.Set("model_default", *patch.ModelDefault)
|
||||
}
|
||||
if patch.Config != nil {
|
||||
b.SetJSON("config", patch.Config)
|
||||
}
|
||||
if patch.Headers != nil {
|
||||
b.SetJSON("headers", patch.Headers)
|
||||
}
|
||||
if patch.Settings != nil {
|
||||
b.SetJSON("settings", patch.Settings)
|
||||
}
|
||||
if patch.IsActive != nil {
|
||||
b.Set("is_active", *patch.IsActive)
|
||||
}
|
||||
if patch.IsPrivate != nil {
|
||||
b.Set("is_private", *patch.IsPrivate)
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ProviderStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM provider_configs WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ProviderStore) ListGlobal(ctx context.Context) ([]models.ProviderConfig, error) {
|
||||
return s.listByScope(ctx, models.ScopeGlobal, "")
|
||||
}
|
||||
|
||||
func (s *ProviderStore) ListForTeam(ctx context.Context, teamID string) ([]models.ProviderConfig, error) {
|
||||
return s.listByScope(ctx, models.ScopeTeam, teamID)
|
||||
}
|
||||
|
||||
func (s *ProviderStore) ListForUser(ctx context.Context, userID string) ([]models.ProviderConfig, error) {
|
||||
return s.listByScope(ctx, models.ScopePersonal, userID)
|
||||
}
|
||||
|
||||
// ListAccessible returns all provider configs a user can access:
|
||||
// global + team (for teams they belong to) + personal.
|
||||
func (s *ProviderStore) ListAccessible(ctx context.Context, userID string) ([]models.ProviderConfig, error) {
|
||||
rows, err := DB.QueryContext(ctx, fmt.Sprintf(`
|
||||
SELECT %s FROM provider_configs
|
||||
WHERE is_active = true AND (
|
||||
scope = 'global'
|
||||
OR (scope = 'personal' AND owner_id = $1)
|
||||
OR (scope = 'team' AND owner_id IN (
|
||||
SELECT team_id FROM team_members WHERE user_id = $1
|
||||
))
|
||||
)
|
||||
ORDER BY scope, name`, providerCols), userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanProviders(rows)
|
||||
}
|
||||
|
||||
// UserCanAccess checks if a user can use a specific provider config.
|
||||
func (s *ProviderStore) UserCanAccess(ctx context.Context, userID, configID string) (bool, error) {
|
||||
var exists bool
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT EXISTS(
|
||||
SELECT 1 FROM provider_configs
|
||||
WHERE id = $2 AND is_active = true AND (
|
||||
scope = 'global'
|
||||
OR (scope = 'personal' AND owner_id = $1)
|
||||
OR (scope = 'team' AND owner_id IN (
|
||||
SELECT team_id FROM team_members WHERE user_id = $1
|
||||
))
|
||||
)
|
||||
)`, userID, configID).Scan(&exists)
|
||||
return exists, err
|
||||
}
|
||||
|
||||
// ── Internal helpers ────────────────────────
|
||||
|
||||
func (s *ProviderStore) listByScope(ctx context.Context, scope, ownerID string) ([]models.ProviderConfig, error) {
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
if scope == models.ScopeGlobal {
|
||||
rows, err = DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM provider_configs WHERE scope = $1 AND is_active = true ORDER BY name", providerCols),
|
||||
scope)
|
||||
} else {
|
||||
rows, err = DB.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT %s FROM provider_configs WHERE scope = $1 AND owner_id = $2 AND is_active = true ORDER BY name", providerCols),
|
||||
scope, ownerID)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
return scanProviders(rows)
|
||||
}
|
||||
|
||||
func scanProvider(row *sql.Row) (*models.ProviderConfig, error) {
|
||||
var p models.ProviderConfig
|
||||
var ownerID, modelDefault sql.NullString
|
||||
var configJSON, headersJSON, settingsJSON []byte
|
||||
err := row.Scan(
|
||||
&p.ID, &p.Scope, &ownerID, &p.Name, &p.Provider, &p.Endpoint,
|
||||
&p.APIKeyEnc, &modelDefault,
|
||||
&configJSON, &headersJSON, &settingsJSON,
|
||||
&p.IsActive, &p.IsPrivate, &p.CreatedAt, &p.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.OwnerID = NullableStringPtr(ownerID)
|
||||
p.ModelDefault = modelDefault.String
|
||||
json.Unmarshal(configJSON, &p.Config)
|
||||
json.Unmarshal(headersJSON, &p.Headers)
|
||||
json.Unmarshal(settingsJSON, &p.Settings)
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func scanProviders(rows *sql.Rows) ([]models.ProviderConfig, error) {
|
||||
var result []models.ProviderConfig
|
||||
for rows.Next() {
|
||||
var p models.ProviderConfig
|
||||
var ownerID, modelDefault sql.NullString
|
||||
var configJSON, headersJSON, settingsJSON []byte
|
||||
err := rows.Scan(
|
||||
&p.ID, &p.Scope, &ownerID, &p.Name, &p.Provider, &p.Endpoint,
|
||||
&p.APIKeyEnc, &modelDefault,
|
||||
&configJSON, &headersJSON, &settingsJSON,
|
||||
&p.IsActive, &p.IsPrivate, &p.CreatedAt, &p.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.OwnerID = NullableStringPtr(ownerID)
|
||||
p.ModelDefault = modelDefault.String
|
||||
json.Unmarshal(configJSON, &p.Config)
|
||||
json.Unmarshal(headersJSON, &p.Headers)
|
||||
json.Unmarshal(settingsJSON, &p.Settings)
|
||||
result = append(result, p)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
27
server/store/postgres/stores.go
Normal file
27
server/store/postgres/stores.go
Normal file
@@ -0,0 +1,27 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
// NewStores creates all Postgres store implementations and wires them
|
||||
// into the Stores bundle. Call this at startup after database.Connect().
|
||||
func NewStores(db *sql.DB) store.Stores {
|
||||
SetDB(db)
|
||||
return store.Stores{
|
||||
Providers: NewProviderStore(),
|
||||
Catalog: NewCatalogStore(),
|
||||
Personas: NewPersonaStore(),
|
||||
Policies: NewPolicyStore(),
|
||||
UserSettings: NewUserModelSettingsStore(),
|
||||
Users: NewUserStore(),
|
||||
Teams: NewTeamStore(),
|
||||
Channels: NewChannelStore(),
|
||||
Messages: NewMessageStore(),
|
||||
Audit: NewAuditStore(),
|
||||
Notes: NewNoteStore(),
|
||||
GlobalConfig: NewGlobalConfigStore(),
|
||||
}
|
||||
}
|
||||
198
server/store/postgres/team.go
Normal file
198
server/store/postgres/team.go
Normal file
@@ -0,0 +1,198 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type TeamStore struct{}
|
||||
|
||||
func NewTeamStore() *TeamStore { return &TeamStore{} }
|
||||
|
||||
func (s *TeamStore) Create(ctx context.Context, t *models.Team) error {
|
||||
settingsJSON := ToJSON(t.Settings)
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO teams (name, description, created_by, is_active, settings)
|
||||
VALUES ($1, $2, $3, $4, $5)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
t.Name, t.Description, t.CreatedBy, t.IsActive, settingsJSON,
|
||||
).Scan(&t.ID, &t.CreatedAt, &t.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *TeamStore) GetByID(ctx context.Context, id string) (*models.Team, error) {
|
||||
var t models.Team
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, name, description, created_by, is_active, settings, created_at, updated_at,
|
||||
(SELECT COUNT(*) FROM team_members WHERE team_id = teams.id) as member_count
|
||||
FROM teams WHERE id = $1`, id).Scan(
|
||||
&t.ID, &t.Name, &t.Description, &t.CreatedBy, &t.IsActive,
|
||||
&settingsJSON, &t.CreatedAt, &t.UpdatedAt, &t.MemberCount,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
json.Unmarshal(settingsJSON, &t.Settings)
|
||||
return &t, nil
|
||||
}
|
||||
|
||||
func (s *TeamStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
b := NewUpdate("teams")
|
||||
for k, v := range fields {
|
||||
if k == "settings" {
|
||||
b.SetJSON(k, v)
|
||||
} else {
|
||||
b.Set(k, v)
|
||||
}
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TeamStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM teams WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TeamStore) List(ctx context.Context) ([]models.Team, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
SELECT id, name, description, created_by, is_active, settings, created_at, updated_at,
|
||||
(SELECT COUNT(*) FROM team_members WHERE team_id = teams.id) as member_count
|
||||
FROM teams ORDER BY name`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.Team
|
||||
for rows.Next() {
|
||||
var t models.Team
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&t.ID, &t.Name, &t.Description, &t.CreatedBy, &t.IsActive,
|
||||
&settingsJSON, &t.CreatedAt, &t.UpdatedAt, &t.MemberCount)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
json.Unmarshal(settingsJSON, &t.Settings)
|
||||
result = append(result, t)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
// ── Members ─────────────────────────────────
|
||||
|
||||
func (s *TeamStore) AddMember(ctx context.Context, teamID, userID, role string) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO team_members (team_id, user_id, role) VALUES ($1, $2, $3)
|
||||
ON CONFLICT (team_id, user_id) DO UPDATE SET role = $3`,
|
||||
teamID, userID, role)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TeamStore) RemoveMember(ctx context.Context, teamID, userID string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM team_members WHERE team_id = $1 AND user_id = $2",
|
||||
teamID, userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TeamStore) UpdateMemberRole(ctx context.Context, teamID, userID, role string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE team_members SET role = $1 WHERE team_id = $2 AND user_id = $3",
|
||||
role, teamID, userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TeamStore) ListMembers(ctx context.Context, teamID string) ([]models.TeamMember, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
SELECT tm.id, tm.team_id, tm.user_id, tm.role, tm.joined_at,
|
||||
u.email, COALESCE(u.display_name, ''), u.username, u.role as user_role
|
||||
FROM team_members tm
|
||||
JOIN users u ON u.id = tm.user_id
|
||||
WHERE tm.team_id = $1
|
||||
ORDER BY tm.role DESC, u.username`, teamID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.TeamMember
|
||||
for rows.Next() {
|
||||
var m models.TeamMember
|
||||
err := rows.Scan(&m.ID, &m.TeamID, &m.UserID, &m.Role, &m.JoinedAt,
|
||||
&m.Email, &m.DisplayName, &m.Username, &m.UserRole)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, m)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (s *TeamStore) GetMember(ctx context.Context, teamID, userID string) (*models.TeamMember, error) {
|
||||
var m models.TeamMember
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT tm.id, tm.team_id, tm.user_id, tm.role, tm.joined_at,
|
||||
u.email, COALESCE(u.display_name, ''), u.username, u.role as user_role
|
||||
FROM team_members tm
|
||||
JOIN users u ON u.id = tm.user_id
|
||||
WHERE tm.team_id = $1 AND tm.user_id = $2`, teamID, userID).Scan(
|
||||
&m.ID, &m.TeamID, &m.UserID, &m.Role, &m.JoinedAt,
|
||||
&m.Email, &m.DisplayName, &m.Username, &m.UserRole)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
// GetUserTeamIDs returns all team IDs a user belongs to.
|
||||
func (s *TeamStore) GetUserTeamIDs(ctx context.Context, userID string) ([]string, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
"SELECT team_id FROM team_members WHERE user_id = $1", userID)
|
||||
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)
|
||||
}
|
||||
return ids, rows.Err()
|
||||
}
|
||||
|
||||
func (s *TeamStore) IsTeamAdmin(ctx context.Context, teamID, userID string) (bool, error) {
|
||||
var role string
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT role FROM team_members WHERE team_id = $1 AND user_id = $2",
|
||||
teamID, userID).Scan(&role)
|
||||
if err == sql.ErrNoRows {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return role == models.TeamRoleAdmin, nil
|
||||
}
|
||||
|
||||
func (s *TeamStore) IsMember(ctx context.Context, teamID, userID string) (bool, error) {
|
||||
var exists bool
|
||||
err := DB.QueryRowContext(ctx,
|
||||
"SELECT EXISTS(SELECT 1 FROM team_members WHERE team_id = $1 AND user_id = $2)",
|
||||
teamID, userID).Scan(&exists)
|
||||
return exists, err
|
||||
}
|
||||
|
||||
// unused but keeping for reference
|
||||
var _ = fmt.Sprintf
|
||||
185
server/store/postgres/user.go
Normal file
185
server/store/postgres/user.go
Normal file
@@ -0,0 +1,185 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type UserStore struct{}
|
||||
|
||||
func NewUserStore() *UserStore { return &UserStore{} }
|
||||
|
||||
func (s *UserStore) Create(ctx context.Context, u *models.User) error {
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO users (username, email, password_hash, display_name, role, is_active, settings)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
u.Username, u.Email, u.PasswordHash, u.DisplayName, u.Role, u.IsActive, ToJSON(u.Settings),
|
||||
).Scan(&u.ID, &u.CreatedAt, &u.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByID(ctx context.Context, id string) (*models.User, error) {
|
||||
return s.getBy(ctx, "id", id)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByUsername(ctx context.Context, username string) (*models.User, error) {
|
||||
return s.getBy(ctx, "username", username)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByEmail(ctx context.Context, email string) (*models.User, error) {
|
||||
return s.getBy(ctx, "email", email)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByLogin(ctx context.Context, login string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE username = $1 OR email = $1`, login).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func (s *UserStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
b := NewUpdate("users")
|
||||
for k, v := range fields {
|
||||
if k == "settings" {
|
||||
b.SetJSON(k, v)
|
||||
} else {
|
||||
b.Set(k, v)
|
||||
}
|
||||
}
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
b.Where("id", id)
|
||||
_, err := b.Exec(DB)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) Delete(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) List(ctx context.Context, opts store.ListOptions) ([]models.User, int, error) {
|
||||
b := NewSelect(
|
||||
"id, username, email, display_name, avatar_url, role, is_active, settings, created_at, updated_at, last_login_at",
|
||||
"users",
|
||||
)
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("username", "ASC")
|
||||
}
|
||||
b.Paginate(opts)
|
||||
|
||||
// Count
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&total)
|
||||
|
||||
q, args := b.Build()
|
||||
rows, err := DB.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.User
|
||||
for rows.Next() {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&u.ID, &u.Username, &u.Email, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
result = append(result, u)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
}
|
||||
|
||||
func (s *UserStore) UpdateLastLogin(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET last_login_at = NOW() WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) SetActive(ctx context.Context, id string, active bool) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET is_active = $1 WHERE id = $2", active, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Refresh Tokens ──────────────────────────
|
||||
|
||||
func (s *UserStore) CreateRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO refresh_tokens (user_id, token_hash, expires_at)
|
||||
VALUES ($1, $2, $3)`, userID, tokenHash, expiresAt)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) GetRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
var userID string
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT user_id FROM refresh_tokens
|
||||
WHERE token_hash = $1 AND revoked_at IS NULL AND expires_at > NOW()`,
|
||||
tokenHash).Scan(&userID)
|
||||
return userID, err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = NOW() WHERE token_hash = $1", tokenHash)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeAllRefreshTokens(ctx context.Context, userID string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = NOW() WHERE user_id = $1 AND revoked_at IS NULL", userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) CleanExpiredTokens(ctx context.Context) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"DELETE FROM refresh_tokens WHERE expires_at < NOW() - INTERVAL '30 days'")
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Internal ────────────────────────────────
|
||||
|
||||
func (s *UserStore) getBy(ctx context.Context, col, val string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE %s = $1`, col), val).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
}
|
||||
163
server/store/postgres/user_settings.go
Normal file
163
server/store/postgres/user_settings.go
Normal file
@@ -0,0 +1,163 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
)
|
||||
|
||||
type UserModelSettingsStore struct{}
|
||||
|
||||
func NewUserModelSettingsStore() *UserModelSettingsStore { return &UserModelSettingsStore{} }
|
||||
|
||||
func (s *UserModelSettingsStore) GetForUser(ctx context.Context, userID string) ([]models.UserModelSetting, error) {
|
||||
rows, err := DB.QueryContext(ctx, `
|
||||
SELECT id, user_id, model_id, hidden, preferred_temperature, preferred_max_tokens,
|
||||
sort_order, created_at, updated_at
|
||||
FROM user_model_settings WHERE user_id = $1 ORDER BY sort_order, model_id`, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var result []models.UserModelSetting
|
||||
for rows.Next() {
|
||||
var s models.UserModelSetting
|
||||
var prefTemp, prefMaxTokens interface{}
|
||||
err := rows.Scan(
|
||||
&s.ID, &s.UserID, &s.ModelID, &s.Hidden,
|
||||
&prefTemp, &prefMaxTokens,
|
||||
&s.SortOrder, &s.CreatedAt, &s.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if f, ok := prefTemp.(float64); ok {
|
||||
s.PreferredTemperature = &f
|
||||
}
|
||||
if n, ok := prefMaxTokens.(int64); ok {
|
||||
v := int(n)
|
||||
s.PreferredMaxTokens = &v
|
||||
}
|
||||
result = append(result, s)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
// GetHiddenModelIDs returns a map of model_id → true for all hidden models.
|
||||
func (s *UserModelSettingsStore) GetHiddenModelIDs(ctx context.Context, userID string) (map[string]bool, error) {
|
||||
rows, err := DB.QueryContext(ctx,
|
||||
"SELECT model_id FROM user_model_settings WHERE user_id = $1 AND hidden = true",
|
||||
userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
result := make(map[string]bool)
|
||||
for rows.Next() {
|
||||
var modelID string
|
||||
if err := rows.Scan(&modelID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[modelID] = true
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
// Set upserts a single user model setting.
|
||||
func (s *UserModelSettingsStore) Set(ctx context.Context, userID, modelID string, patch models.UserModelSettingPatch) error {
|
||||
b := NewUpdate("user_model_settings")
|
||||
if patch.Hidden != nil {
|
||||
b.Set("hidden", *patch.Hidden)
|
||||
}
|
||||
if patch.PreferredTemperature != nil {
|
||||
b.Set("preferred_temperature", models.NullFloat(patch.PreferredTemperature))
|
||||
}
|
||||
if patch.PreferredMaxTokens != nil {
|
||||
b.Set("preferred_max_tokens", models.NullInt(patch.PreferredMaxTokens))
|
||||
}
|
||||
if patch.SortOrder != nil {
|
||||
b.Set("sort_order", *patch.SortOrder)
|
||||
}
|
||||
|
||||
if !b.HasSets() {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Use upsert: insert if not exists, update if exists
|
||||
// Build a custom upsert since the update builder doesn't handle INSERT ON CONFLICT
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO user_model_settings (user_id, model_id, hidden, preferred_temperature, preferred_max_tokens, sort_order)
|
||||
VALUES ($1, $2, $3, $4, $5, $6)
|
||||
ON CONFLICT (user_id, model_id)
|
||||
DO UPDATE SET
|
||||
hidden = COALESCE($3, user_model_settings.hidden),
|
||||
preferred_temperature = COALESCE($4, user_model_settings.preferred_temperature),
|
||||
preferred_max_tokens = COALESCE($5, user_model_settings.preferred_max_tokens),
|
||||
sort_order = COALESCE($6, user_model_settings.sort_order)`,
|
||||
userID, modelID,
|
||||
patchBoolOrNil(patch.Hidden),
|
||||
patchFloat64OrNil(patch.PreferredTemperature),
|
||||
patchIntOrNil(patch.PreferredMaxTokens),
|
||||
patchIntOrNil(patch.SortOrder),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
// BulkSetHidden sets the hidden state for multiple models at once.
|
||||
func (s *UserModelSettingsStore) BulkSetHidden(ctx context.Context, userID string, modelIDs []string, hidden bool) error {
|
||||
if len(modelIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build parameterized IN clause
|
||||
placeholders := make([]string, len(modelIDs))
|
||||
args := make([]interface{}, 0, len(modelIDs)+2)
|
||||
args = append(args, userID, hidden)
|
||||
for i, id := range modelIDs {
|
||||
placeholders[i] = fmt.Sprintf("$%d", i+3)
|
||||
args = append(args, id)
|
||||
}
|
||||
|
||||
// Upsert each: some might not have rows yet
|
||||
for _, modelID := range modelIDs {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO user_model_settings (user_id, model_id, hidden)
|
||||
VALUES ($1, $2, $3)
|
||||
ON CONFLICT (user_id, model_id) DO UPDATE SET hidden = $3`,
|
||||
userID, modelID, hidden)
|
||||
if err != nil {
|
||||
return fmt.Errorf("set hidden for %s: %w", modelID, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── Helpers ─────────────────────────────────
|
||||
|
||||
func patchBoolOrNil(b *bool) interface{} {
|
||||
if b == nil {
|
||||
return nil
|
||||
}
|
||||
return *b
|
||||
}
|
||||
|
||||
func patchFloat64OrNil(f *float64) interface{} {
|
||||
if f == nil {
|
||||
return nil
|
||||
}
|
||||
return *f
|
||||
}
|
||||
|
||||
func patchIntOrNil(i *int) interface{} {
|
||||
if i == nil {
|
||||
return nil
|
||||
}
|
||||
return *i
|
||||
}
|
||||
|
||||
// unused but keeping for reference - will be used in ListOptions-based queries
|
||||
var _ = strings.Join
|
||||
Reference in New Issue
Block a user