Files
descrybe/apps/api/internal/aiprovider/service.go
T
2026-08-23 22:03:57 +02:00

479 lines
15 KiB
Go

package aiprovider
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/descrybe/descrybe-v2/apps/api/internal/aiaudit"
"github.com/descrybe/descrybe-v2/apps/api/internal/platformsettings"
"github.com/descrybe/descrybe-v2/apps/api/internal/processing"
"github.com/descrybe/descrybe-v2/apps/api/internal/security"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
var (
ErrNotConfigured = errors.New("ai provider not configured")
ErrInvalidMode = errors.New("mode must be internal, popular, or custom")
ErrInvalidPopular = errors.New("unknown popular provider")
ErrMissingAPIKey = errors.New("api key required")
ErrMissingModel = errors.New("model required")
ErrMissingURL = errors.New("base url required for custom provider")
)
// aiProbeTimeout bounds admin/company connection tests so a hung provider cannot
// hold the HTTP request for multi-retry OpenAI client durations.
const aiProbeTimeout = 45 * time.Second
type Service struct {
Pool *pgxpool.Pool
Key []byte
Env EnvConfig
// Platform is optional; when set, platform OpenAI is loaded from admin
// settings (DB) with EnvConfig as bootstrap fallback.
Platform *platformsettings.Service
// Roles is optional admin role-binding lookup (processing / embeddings / …).
// When nil or a role is unset, ResolveCompleterForRole falls back to Resolve.
Roles RoleEndpointSource
HTTPClient *http.Client
// Audit optionally captures every resolved Completer's prompts/responses for
// the admin AI inspector. Nil disables capture (see audit.go).
Audit *aiaudit.Recorder
}
func NewService(pool *pgxpool.Pool, env EnvConfig) *Service {
keyMaterial := firstNonEmpty(env.AppEncryptionKey, env.CredentialsEncryptionKey, env.TokenSigningSecret)
return &Service{
Pool: pool,
Key: DeriveKey(keyMaterial, env.DatabaseURL),
Env: env,
// HTTPClient is optional (tests). Production uses NewOpenAIClient's
// SafeHTTPClient so dial-time SSRF applies; leave nil here so platform
// OPENAI_BASE_URL loopback (local models) is not overwritten.
}
}
type stored struct {
mode, popularName, baseURL, model, keyEnc, last4 string
enabled bool
lastTest *time.Time
lastStatus *string
}
func (s *Service) loadStored(ctx context.Context, companyID uuid.UUID) (stored, error) {
var st stored
err := s.Pool.QueryRow(ctx, `
SELECT mode, popular_name, base_url, model, api_key_enc, api_key_last4, is_enabled,
last_test_at, last_test_status
FROM ai_providers WHERE company_id = $1`, companyID).Scan(
&st.mode, &st.popularName, &st.baseURL, &st.model, &st.keyEnc, &st.last4, &st.enabled,
&st.lastTest, &st.lastStatus,
)
return st, err
}
func (s *Service) GetConfig(ctx context.Context, companyID uuid.UUID) (PublicConfig, error) {
platformOK, err := s.platformConfigured(ctx)
if err != nil {
return PublicConfig{}, err
}
st, err := s.loadStored(ctx, companyID)
if errors.Is(err, pgx.ErrNoRows) {
return PublicConfig{
Mode: ModeInternal,
Configured: false,
IsEnabled: false,
ActiveModeLabel: ModeInternalLabel,
PlatformFallback: platformOK,
PopularProviders: PopularCatalog,
}, nil
}
if err != nil {
return PublicConfig{}, err
}
hasKey := st.keyEnc != ""
masked := ""
if hasKey && st.last4 != "" {
masked = "••••" + st.last4
}
active := ModeInternalLabel
if st.enabled && hasKey && (st.mode == ModePopular || st.mode == ModeCustom) {
active = AnalyticsMode(st.mode, st.popularName)
}
return PublicConfig{
Mode: normalizeMode(st.mode),
PopularName: st.popularName,
BaseURL: st.baseURL,
Model: st.model,
IsEnabled: st.enabled,
Configured: true,
HasAPIKey: hasKey,
APIKeyLast4: st.last4,
APIKeyMasked: masked,
LastTestAt: st.lastTest,
LastTestStatus: st.lastStatus,
ActiveModeLabel: active,
PlatformFallback: platformOK,
PopularProviders: PopularCatalog,
}, nil
}
func (s *Service) UpdateConfig(ctx context.Context, companyID uuid.UUID, in UpdateInput) (PublicConfig, error) {
mode := normalizeMode(in.Mode)
if mode != ModeInternal && mode != ModePopular && mode != ModeCustom {
return PublicConfig{}, ErrInvalidMode
}
var existing stored
existing, err := s.loadStored(ctx, companyID)
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
return PublicConfig{}, err
}
hasExisting := err == nil
keyEnc := ""
last4v := ""
if hasExisting {
keyEnc = existing.keyEnc
last4v = existing.last4
}
if in.ClearAPIKey {
keyEnc = ""
last4v = ""
} else if strings.TrimSpace(in.APIKey) != "" {
plain := strings.TrimSpace(in.APIKey)
enc, err := EncryptSecret(s.Key, plain)
if err != nil {
return PublicConfig{}, err
}
keyEnc = enc
last4v = last4(plain)
}
popularName := ""
baseURL := ""
model := strings.TrimSpace(in.Model)
switch mode {
case ModeInternal:
// Platform fallback; company key optional/cleared when switching away from BYOK.
if !in.IsEnabled {
keyEnc = ""
last4v = ""
}
case ModePopular:
pop, ok := FindPopular(in.PopularName)
if !ok {
return PublicConfig{}, ErrInvalidPopular
}
popularName = pop.Name
baseURL = pop.BaseURL
if model == "" {
model = pop.DefaultModel
}
if !modelAllowed(pop, model) {
return PublicConfig{}, ClientMsg(fmt.Sprintf("model %q is not in the %s catalog (or leave blank for default)", model, pop.Name))
}
if in.IsEnabled && keyEnc == "" {
return PublicConfig{}, ErrMissingAPIKey
}
case ModeCustom:
normalized, err := validateProviderBaseURL(in.BaseURL)
if err != nil {
return PublicConfig{}, err
}
baseURL = normalized
if model == "" {
return PublicConfig{}, ErrMissingModel
}
if in.IsEnabled && keyEnc == "" {
return PublicConfig{}, ErrMissingAPIKey
}
}
_, err = s.Pool.Exec(ctx, `
INSERT INTO ai_providers (
company_id, mode, popular_name, base_url, model, api_key_enc, api_key_last4, is_enabled, updated_at
) VALUES ($1,$2,$3,$4,$5,$6,$7,$8, now())
ON CONFLICT (company_id) DO UPDATE SET
mode = EXCLUDED.mode,
popular_name = EXCLUDED.popular_name,
base_url = EXCLUDED.base_url,
model = EXCLUDED.model,
api_key_enc = EXCLUDED.api_key_enc,
api_key_last4 = EXCLUDED.api_key_last4,
is_enabled = EXCLUDED.is_enabled,
updated_at = now()`,
companyID, mode, popularName, baseURL, model, keyEnc, last4v, in.IsEnabled && mode != ModeInternal)
if err != nil {
return PublicConfig{}, err
}
return s.GetConfig(ctx, companyID)
}
func modelAllowed(pop PopularProvider, model string) bool {
model = strings.TrimSpace(model)
if model == "" || model == pop.DefaultModel {
return true
}
for _, m := range pop.Models {
if m == model {
return true
}
}
// Allow unknown model strings for popular providers (API may add models faster than catalog).
return true
}
func validateProviderBaseURL(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", ErrMissingURL
}
normalized, err := security.ValidatePublicHTTPSURL(raw)
if err != nil {
return "", err
}
if normalized == "" {
return "", ErrMissingURL
}
return strings.TrimRight(normalized, "/"), nil
}
// Resolve picks company BYOK completer when enabled+keyed, else platform AI
// from admin settings (DB), with optional env fallback via Platform / Env.
func (s *Service) Resolve(ctx context.Context, companyID uuid.UUID) (Resolved, error) {
rpm := s.Env.ProcessingRPM
retries := s.Env.ProcessingMaxRetries
st, err := s.loadStored(ctx, companyID)
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
return Resolved{}, err
}
if err == nil && st.enabled && (st.mode == ModePopular || st.mode == ModeCustom) {
key, derr := DecryptSecret(s.Key, st.keyEnc)
if derr != nil {
return Resolved{}, derr
}
if strings.TrimSpace(key) != "" && strings.TrimSpace(st.baseURL) != "" && strings.TrimSpace(st.model) != "" {
client := processing.NewOpenAIClient(key, st.baseURL, st.model, rpm, retries)
client.ModeLabel = AnalyticsMode(st.mode, st.popularName)
if s.HTTPClient != nil {
client.HTTPClient = s.HTTPClient
}
return Resolved{
Completer: client,
ModeLabel: client.ModeLabel,
UsingBYOK: true,
}, nil
}
}
platform, err := s.resolvePlatformOpenAI(ctx)
if err != nil {
return Resolved{}, err
}
if strings.TrimSpace(platform.APIKey) == "" {
return Resolved{ModeLabel: ModeInternalLabel, UsingBYOK: false}, nil
}
modeLabel := platformModeLabel(platform.Source, platform.BaseURL)
client := processing.NewOpenAIClient(
platform.APIKey,
platform.BaseURL,
platform.Model,
rpm,
retries,
)
client.ModeLabel = modeLabel
if s.HTTPClient != nil {
client.HTTPClient = s.HTTPClient
}
return Resolved{
Completer: client,
ModeLabel: modeLabel,
UsingBYOK: false,
}, nil
}
// platformModeLabel maps platform OpenAI source/base to analytics ModeLabel.
// DB-backed admin settings → internal; env bootstrap or mock/loopback → custom.
func platformModeLabel(source, baseURL string) string {
if processing.IsMockOrLoopbackBaseURL(baseURL) {
return ModeCustomLabel
}
if strings.TrimSpace(source) == platformsettings.SourceEnv {
return ModeCustomLabel
}
return ModeInternalLabel
}
func (s *Service) platformConfigured(ctx context.Context) (bool, error) {
oi, err := s.resolvePlatformOpenAI(ctx)
if err != nil {
return false, err
}
return strings.TrimSpace(oi.APIKey) != "", nil
}
func (s *Service) resolvePlatformOpenAI(ctx context.Context) (platformsettings.ResolvedOpenAI, error) {
if s.Platform != nil {
return s.Platform.ResolveOpenAI(ctx)
}
out := platformsettings.ResolvedOpenAI{
APIKey: strings.TrimSpace(s.Env.OpenAIAPIKey),
BaseURL: strings.TrimSpace(s.Env.OpenAIBaseURL),
Model: strings.TrimSpace(s.Env.OpenAIModel),
Source: platformsettings.SourceNone,
}
if out.APIKey != "" {
out.Source = platformsettings.SourceEnv
}
return out, nil
}
// ResolveCompleter implements processing.CompanyCompleterResolver (legacy callers).
// Prefer ResolveCompleterForRole for new call sites.
func (s *Service) ResolveCompleter(ctx context.Context, companyID uuid.UUID) (processing.Completer, string, bool, error) {
r, err := s.Resolve(ctx, companyID)
if err != nil {
return nil, ModeInternalLabel, false, err
}
return s.wrapAudit(r.Completer, companyID, RoleProcessing, r.ModeLabel), r.ModeLabel, r.UsingBYOK, nil
}
// TestPlatformRole probes admin platform AI role credentials (not company BYOK).
// Chat roles send a minimal completion; vectorization sends a one-token embed.
// Never returns upstream error bodies (may contain key fragments).
func (s *Service) TestPlatformRole(ctx context.Context, role string) (map[string]any, error) {
ctx, cancel := context.WithTimeout(ctx, aiProbeTimeout)
defer cancel()
role = strings.TrimSpace(role)
out := map[string]any{"role": role}
if role == "" || !platformsettings.ValidAIRole(role) {
out["status"] = "failed"
out["message"] = "unknown ai role"
return out, fmt.Errorf("unknown ai role %q", role)
}
if role == RoleVectorization {
emb, err := s.ResolveEmbedderForRole(ctx, uuid.Nil, role)
if err != nil {
out["status"] = "failed"
out["message"] = "provider resolve failed"
return out, err
}
if emb == nil {
out["status"] = "skipped"
out["message"] = "Vectorization AI is not configured in admin platform settings"
return out, nil
}
if _, err := emb.Embed(ctx, []string{"ping"}); err != nil {
out["status"] = "failed"
out["message"] = "connection failed — check vectorization provider, key, and model"
return out, err
}
out["status"] = "ok"
out["message"] = "Embeddings probe succeeded"
return out, nil
}
completer, _, _, err := s.resolvePlatformRoleCompleter(ctx, role)
if err != nil {
out["status"] = "failed"
out["message"] = "provider resolve failed"
return out, err
}
if completer == nil {
out["status"] = "skipped"
out["message"] = "AI role is not configured (or disabled) in admin platform settings"
return out, nil
}
if err := probeCompleter(ctx, completer); err != nil {
out["status"] = "failed"
// Prefer classified probe detail (no secrets / model-id dumps) over a generic blur.
msg := processing.TruncateError(err)
if msg == "" || msg == "processing_failed" || msg == "provider error (details redacted)" {
msg = "connection failed — check provider, key, base URL, and model"
}
out["message"] = msg
return out, err
}
out["status"] = "ok"
out["message"] = "Connection probe succeeded"
return out, nil
}
// probeCompleter runs a minimal chat completion with GPT-5.6-safe options
// (max_completion_tokens + reasoning_effort=none, no custom temperature).
func probeCompleter(ctx context.Context, completer processing.Completer) error {
if completer == nil {
return ErrNotConfigured
}
opts := processing.ProbeCompleteOptions()
if co, ok := completer.(processing.CompleterWithOptions); ok {
_, err := co.CompleteWithOptions(ctx, "Reply with exactly: ok", "ping", opts)
return err
}
_, err := completer.Complete(ctx, "Reply with exactly: ok", "ping")
return err
}
// TestConnection sends a minimal chat completion and records last_test_*.
func (s *Service) TestConnection(ctx context.Context, companyID uuid.UUID) (map[string]any, error) {
ctx, cancel := context.WithTimeout(ctx, aiProbeTimeout)
defer cancel()
resolved, err := s.Resolve(ctx, companyID)
status := "ok"
message := "connection successful"
if err != nil {
status = "failed"
message = "provider resolve failed"
_, _ = s.Pool.Exec(ctx, `
UPDATE ai_providers SET last_test_at = now(), last_test_status = $2, updated_at = now()
WHERE company_id = $1`, companyID, status)
return map[string]any{"status": status, "message": message, "mode": ModeInternalLabel}, err
}
if resolved.Completer == nil {
status = "failed"
message = "no api key configured (company BYOK or admin platform settings)"
_, _ = s.Pool.Exec(ctx, `
UPDATE ai_providers SET last_test_at = now(), last_test_status = $2, updated_at = now()
WHERE company_id = $1`, companyID, status)
return map[string]any{"status": status, "message": message, "mode": resolved.ModeLabel}, ErrNotConfigured
}
if err := probeCompleter(ctx, resolved.Completer); err != nil {
status = "failed"
message = processing.TruncateError(err)
if message == "" || message == "processing_failed" || message == "provider error (details redacted)" {
message = "connection failed — check provider, key, base URL, and model"
}
_, _ = s.Pool.Exec(ctx, `
UPDATE ai_providers SET last_test_at = now(), last_test_status = $2, updated_at = now()
WHERE company_id = $1`, companyID, status)
return map[string]any{"status": status, "message": message, "mode": resolved.ModeLabel}, err
}
_, _ = s.Pool.Exec(ctx, `
UPDATE ai_providers SET last_test_at = now(), last_test_status = $2, updated_at = now()
WHERE company_id = $1`, companyID, status)
return map[string]any{
"status": status,
"message": message,
"mode": resolved.ModeLabel,
"byok": resolved.UsingBYOK,
}, nil
}
func firstNonEmpty(vals ...string) string {
for _, v := range vals {
if strings.TrimSpace(v) != "" {
return v
}
}
return ""
}