86 lines
2.7 KiB
Go
86 lines
2.7 KiB
Go
package postgres
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
|
|
"git.gobha.me/xcaliber/chat-switchboard/models"
|
|
)
|
|
|
|
// CapOverrideStore implements store.CapabilityOverrideStore for Postgres.
|
|
type CapOverrideStore struct {
|
|
db *sql.DB
|
|
}
|
|
|
|
func NewCapOverrideStore(db *sql.DB) *CapOverrideStore {
|
|
return &CapOverrideStore{db: db}
|
|
}
|
|
|
|
func (s *CapOverrideStore) Set(ctx context.Context, o *models.CapabilityOverride) error {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
INSERT INTO capability_overrides (provider_config_id, model_id, field, value, set_by)
|
|
VALUES ($1, $2, $3, $4, $5)
|
|
ON CONFLICT (provider_config_id, model_id, field) DO UPDATE SET
|
|
value = EXCLUDED.value,
|
|
set_by = EXCLUDED.set_by,
|
|
created_at = now()
|
|
`, o.ProviderConfigID, o.ModelID, o.Field, o.Value, o.SetBy)
|
|
return err
|
|
}
|
|
|
|
func (s *CapOverrideStore) Delete(ctx context.Context, id string) error {
|
|
_, err := s.db.ExecContext(ctx, `DELETE FROM capability_overrides WHERE id = $1`, id)
|
|
return err
|
|
}
|
|
|
|
func (s *CapOverrideStore) ListForModel(ctx context.Context, modelID string) ([]models.CapabilityOverride, error) {
|
|
return s.query(ctx, `
|
|
SELECT id, provider_config_id, model_id, field, value, set_by, created_at
|
|
FROM capability_overrides
|
|
WHERE model_id = $1
|
|
ORDER BY provider_config_id NULLS LAST, field
|
|
`, modelID)
|
|
}
|
|
|
|
func (s *CapOverrideStore) ListForProviderModel(ctx context.Context, providerConfigID, modelID string) ([]models.CapabilityOverride, error) {
|
|
return s.query(ctx, `
|
|
SELECT id, provider_config_id, model_id, field, value, set_by, created_at
|
|
FROM capability_overrides
|
|
WHERE model_id = $1 AND (provider_config_id = $2 OR provider_config_id IS NULL)
|
|
ORDER BY provider_config_id NULLS LAST, field
|
|
`, modelID, providerConfigID)
|
|
}
|
|
|
|
func (s *CapOverrideStore) ListAll(ctx context.Context) ([]models.CapabilityOverride, error) {
|
|
return s.query(ctx, `
|
|
SELECT id, provider_config_id, model_id, field, value, set_by, created_at
|
|
FROM capability_overrides
|
|
ORDER BY model_id, provider_config_id NULLS LAST, field
|
|
`)
|
|
}
|
|
|
|
func (s *CapOverrideStore) DeleteForProvider(ctx context.Context, providerConfigID string) error {
|
|
_, err := s.db.ExecContext(ctx, `
|
|
DELETE FROM capability_overrides WHERE provider_config_id = $1
|
|
`, providerConfigID)
|
|
return err
|
|
}
|
|
|
|
func (s *CapOverrideStore) query(ctx context.Context, q string, args ...interface{}) ([]models.CapabilityOverride, error) {
|
|
rows, err := s.db.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
|
|
var result []models.CapabilityOverride
|
|
for rows.Next() {
|
|
var o models.CapabilityOverride
|
|
if err := rows.Scan(&o.ID, &o.ProviderConfigID, &o.ModelID, &o.Field, &o.Value, &o.SetBy, &o.CreatedAt); err != nil {
|
|
return nil, err
|
|
}
|
|
result = append(result, o)
|
|
}
|
|
return result, rows.Err()
|
|
}
|