Files
ResolveSpec/pkg/aiproxy/hooks.go
T
warkanum 63cb0d22eb feat(aiproxy): add authenticated proxy for OpenAI-compatible APIs and MCP servers
Hides upstream URLs and keys behind the security layer. Per-upstream roles,
model/tool allowlists, per-user rate limits, Before/AfterProxy hooks, audit
sink, Prometheus metrics (metrics.AIProxyRecorder) and upstreams loaded from
the resolvespec_ai_proxies stored procedure.
2026-10-07 23:39:40 +02:00

97 lines
2.3 KiB
Go

package aiproxy
import (
"context"
"fmt"
"net/http"
"sync"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/security"
)
// HookType defines when a hook runs.
type HookType string
const (
// BeforeProxy runs after auth, rate limit and request inspection, before the upstream call.
BeforeProxy HookType = "before_proxy"
// AfterProxy runs once the response finished (or failed). Errors are logged only.
AfterProxy HookType = "after_proxy"
)
// Usage is the token usage reported by an OpenAI upstream.
type Usage struct {
PromptTokens int64
CompletionTokens int64
TotalTokens int64
}
// HookContext is passed to hooks.
type HookContext struct {
Context context.Context
Request *http.Request // BeforeProxy may modify headers/query; the key is injected later
UserContext *security.UserContext
Upstream string
Kind Kind
Model string // OpenAI: requested model
Tools []string // MCP: tools named by tools/call
// AfterProxy only
StatusCode int
Duration time.Duration
Usage Usage // OpenAI only; needs stream_options.include_usage for streams
Error error
// BeforeProxy only
Abort bool
AbortMessage string
AbortCode int // default 403
}
// HookFunc is a hook. A returned error aborts a BeforeProxy request.
type HookFunc func(*HookContext) error
// HookRegistry holds hooks per type.
type HookRegistry struct {
mu sync.RWMutex
hooks map[HookType][]HookFunc
}
// NewHookRegistry creates an empty registry.
func NewHookRegistry() *HookRegistry {
return &HookRegistry{hooks: make(map[HookType][]HookFunc)}
}
// Register adds a hook.
func (r *HookRegistry) Register(t HookType, h HookFunc) {
r.mu.Lock()
defer r.mu.Unlock()
r.hooks[t] = append(r.hooks[t], h)
}
// Execute runs the hooks in order and stops at the first error or abort.
func (r *HookRegistry) Execute(t HookType, ctx *HookContext) error {
r.mu.RLock()
list := append([]HookFunc(nil), r.hooks[t]...)
r.mu.RUnlock()
for i, h := range list {
if err := h(ctx); err != nil {
return fmt.Errorf("aiproxy hook %d for %s failed: %w", i+1, t, err)
}
if ctx.Abort {
return fmt.Errorf("aborted by hook: %s", ctx.AbortMessage)
}
}
return nil
}
func (r *HookRegistry) executeAfter(ctx *HookContext) {
if err := r.Execute(AfterProxy, ctx); err != nil {
logger.Warn("aiproxy: %v", err)
}
}