Files
gronod eab16991ac Implement the runnable MCP server: 31 tools, auth, projections (Phase 07)
Lands the executable half of the Phase 01-06 contracts: stdio MCP
wiring with bundle policy enforced at call time, JWT/api_key upstream
auth with single-flight renewal and 401 retry, strict argument
decoding, allowlisted result projections with enum label twins, TV
season expansion, settings read-modify-write under a revision lock,
and a sanitized ToolError envelope that never forwards raw upstream
bodies.
2026-09-18 19:14:10 +01:00

188 lines
5.0 KiB
Go

package ombi
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"sync"
"time"
)
// ErrAuthFailed is returned when the upstream rejects login credentials
// or returns an unusable token response.
var ErrAuthFailed = errors.New("upstream authentication failed")
// AuthMode is the selected upstream credential mechanism.
type AuthMode string
const (
AuthModeJWT AuthMode = "jwt"
AuthModeAPIKey AuthMode = "api_key"
)
// Credentials holds only what the selected mode requires.
type Credentials struct {
Mode AuthMode
Username string // OMBI_USERNAME (jwt)
Password string // OMBI_PASSWORD (jwt)
APIKey string // OMBI_API_KEY (api_key)
UserName string // OMBI_USER_NAME (api_key, optional)
}
// AuthManager supplies per-request upstream credentials and owns the
// JWT lifecycle. Safe for concurrent use.
type AuthManager struct {
creds Credentials
base string
http *http.Client
mu sync.Mutex
token string
expiresAt time.Time
renewing chan struct{} // non-nil while a single-flight renewal runs
}
const expirySkew = 60 * time.Second
const fallbackTokenTTL = time.Hour
// NewAuthManager builds an AuthManager for the given credentials.
// base is the configured Ombi base URL (any path prefix preserved);
// hc is the shared HTTP client used for token requests.
func NewAuthManager(creds Credentials, base string, hc *http.Client) *AuthManager {
return &AuthManager{creds: creds, base: base, http: hc}
}
// Apply sets the authentication headers on an outgoing upstream request.
// In api_key mode it sets ApiKey (+ optional UserName) and returns.
// In jwt mode it blocks until a valid Bearer token is available.
func (a *AuthManager) Apply(ctx context.Context, req *http.Request) error {
if a.creds.Mode == AuthModeAPIKey {
req.Header.Set("ApiKey", a.creds.APIKey)
if a.creds.UserName != "" {
req.Header.Set("UserName", a.creds.UserName)
}
return nil
}
tok, err := a.getToken(ctx)
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+tok)
return nil
}
// Invalidate discards the cached token (called on upstream 401).
func (a *AuthManager) Invalidate() {
a.mu.Lock()
a.token = ""
a.expiresAt = time.Time{}
a.mu.Unlock()
}
// getToken returns a valid token, single-flight renewing if expired.
func (a *AuthManager) getToken(ctx context.Context) (string, error) {
a.mu.Lock()
if a.token != "" && time.Now().Add(expirySkew).Before(a.expiresAt) {
defer a.mu.Unlock()
return a.token, nil
}
if a.renewing != nil { // another caller renews; wait for it
done := a.renewing
a.mu.Unlock()
select {
case <-done:
return a.getToken(ctx)
case <-ctx.Done():
return "", ctx.Err()
}
}
a.renewing = make(chan struct{})
a.mu.Unlock()
defer func() {
a.mu.Lock()
close(a.renewing)
a.renewing = nil
a.mu.Unlock()
}()
return a.login(ctx)
}
// login performs a full POST /api/v1/Token login, caches the token and
// resolves its expiry: RFC3339 expiration field, then the JWT exp claim
// (decode only — not signature validation), else a conservative TTL.
func (a *AuthManager) login(ctx context.Context) (string, error) {
body, err := json.Marshal(map[string]string{
"username": a.creds.Username,
"password": a.creds.Password,
})
if err != nil {
return "", err
}
u := strings.TrimSuffix(a.base, "/") + "/api/v1/Token"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
resp, err := a.http.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
return "", fmt.Errorf("%w: login rejected with HTTP %d", ErrAuthFailed, resp.StatusCode)
}
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return "", fmt.Errorf("token endpoint returned HTTP %d", resp.StatusCode)
}
var tok struct {
AccessToken string `json:"access_token"`
Expiration string `json:"expiration"`
}
if err := json.NewDecoder(resp.Body).Decode(&tok); err != nil {
return "", fmt.Errorf("token response decode: %w", err)
}
if tok.AccessToken == "" {
return "", fmt.Errorf("%w: token response missing access_token", ErrAuthFailed)
}
expiresAt := resolveExpiry(tok.AccessToken, tok.Expiration)
a.mu.Lock()
a.token = tok.AccessToken
a.expiresAt = expiresAt
a.mu.Unlock()
return tok.AccessToken, nil
}
func resolveExpiry(token, expiration string) time.Time {
if t, err := time.Parse(time.RFC3339, expiration); err == nil {
return t
}
if exp, ok := jwtExpClaim(token); ok {
return time.Unix(exp, 0)
}
return time.Now().Add(fallbackTokenTTL)
}
func jwtExpClaim(token string) (int64, bool) {
parts := strings.Split(token, ".")
if len(parts) < 2 {
return 0, false
}
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return 0, false
}
var claims struct {
Exp int64 `json:"exp"`
}
if err := json.Unmarshal(payload, &claims); err != nil || claims.Exp == 0 {
return 0, false
}
return claims.Exp, true
}