This commit is contained in:
2026-09-11 06:14:38 +02:00
parent bf64652300
commit e581949946
161 changed files with 31126 additions and 1 deletions
+492
View File
@@ -0,0 +1,492 @@
package auth
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"net"
"net/http"
"sort"
"strings"
"sync"
"time"
"github.com/example/ollama-fair-gateway/internal/config"
)
type Identity struct {
Tenant string `json:"tenant"`
Subject string `json:"subject"`
Application string `json:"application,omitempty"`
AuthType string `json:"auth_type"`
Scopes map[string]bool `json:"-"`
ClientIP string `json:"client_ip"`
ModelACLSet bool `json:"-"`
ModelAccess config.ModelAccessRule `json:"-"`
ServiceClass string `json:"-"`
}
func (i Identity) Actor() string {
// Interactive OIDC identities are fair-scheduled by subject, so one user
// cannot gain extra shares by using multiple OAuth clients. Static API-key
// and trusted-IP identities are normally applications and use that identity.
if i.AuthType != "oidc" && i.Application != "" {
return "app:" + i.Application
}
if i.Subject != "" {
return i.Subject
}
if i.Application != "" {
return "app:" + i.Application
}
return "anonymous"
}
func (i Identity) HasScope(s string) bool { return i.Scopes[s] || i.Scopes["*"] }
func (i Identity) IsAdmin() bool { return i.HasScope("gateway:admin") }
type ctxKey struct{}
func WithIdentity(ctx context.Context, i Identity) context.Context {
return context.WithValue(ctx, ctxKey{}, i)
}
func FromContext(ctx context.Context) (Identity, bool) {
i, ok := ctx.Value(ctxKey{}).(Identity)
return i, ok
}
// APIKeyInfo is safe to return through the admin API. It never contains the
// key secret. UI-created keys may be backed by a durable RuntimeKeyStore.
type APIKeyInfo struct {
ID string `json:"id,omitempty"`
Name string `json:"name"`
Tenant string `json:"tenant"`
Subject string `json:"subject"`
Application string `json:"application,omitempty"`
Scopes []string `json:"scopes"`
AllowedModels []string `json:"allowed_models,omitempty"`
DeniedModels []string `json:"denied_models,omitempty"`
ServiceClass string `json:"service_class,omitempty"`
Source string `json:"source"` // config | runtime | persistent
KeyHint string `json:"key_hint,omitempty"`
CreatedAt *time.Time `json:"created_at,omitempty"`
Deletable bool `json:"deletable"`
}
type APIKeyCreate struct {
Name string
Tenant string
Subject string
Application string
Scopes []string
AllowedModels []string
DeniedModels []string
ServiceClass string
}
// StoredAPIKey is the durable representation of a UI-created API key. Only
// the SHA-256 hash is persisted; the plaintext secret never leaves CreateAPIKey.
type StoredAPIKey struct {
ID string `json:"id"`
Name string `json:"name"`
Tenant string `json:"tenant"`
Subject string `json:"subject"`
Application string `json:"application,omitempty"`
Scopes []string `json:"scopes"`
AllowedModels []string `json:"allowed_models,omitempty"`
DeniedModels []string `json:"denied_models,omitempty"`
ServiceClass string `json:"service_class,omitempty"`
HashHex string `json:"hash_sha256"`
KeyHint string `json:"key_hint,omitempty"`
CreatedAt time.Time `json:"created_at"`
}
type RuntimeKeyStore interface {
Load() ([]StoredAPIKey, error)
Put(StoredAPIKey) error
Delete(string) error
Health(context.Context) error
}
type managedKey struct {
identity Identity
info APIKeyInfo
}
type bypassRule struct {
nets []*net.IPNet
identity Identity
}
type Authenticator struct {
oidc *OIDCVerifier
mu sync.RWMutex
keys map[[32]byte]managedKey
runtime map[string][32]byte
bypass []bypassRule
trusted []*net.IPNet
bypassForwarded bool
store RuntimeKeyStore
}
func New(ctx context.Context, cfg config.AuthConfig) (*Authenticator, error) {
return NewWithRuntimeStore(ctx, cfg, nil)
}
func NewWithRuntimeStore(ctx context.Context, cfg config.AuthConfig, store RuntimeKeyStore) (*Authenticator, error) {
a := &Authenticator{keys: make(map[[32]byte]managedKey), runtime: make(map[string][32]byte), bypassForwarded: cfg.IPBypassUseForwardedIP, store: store}
for _, c := range cfg.TrustedProxies {
_, n, _ := net.ParseCIDR(c)
a.trusted = append(a.trusted, n)
}
for _, b := range cfg.IPBypass {
r := bypassRule{identity: Identity{Tenant: b.Tenant, Subject: b.Subject, Application: b.Application, AuthType: "ip-bypass", Scopes: scopeMap(b.Scopes)}}
if r.identity.Subject == "" {
r.identity.Subject = "ip-bypass"
}
for _, c := range b.CIDRs {
_, n, _ := net.ParseCIDR(c)
r.nets = append(r.nets, n)
}
a.bypass = append(a.bypass, r)
}
for _, k := range cfg.APIKeys {
if k.Key == "" {
return nil, fmt.Errorf("api key %q is empty (is its environment variable set?)", k.Name)
}
id := Identity{Tenant: k.Tenant, Subject: k.Subject, Application: k.Application, AuthType: "api-key", Scopes: scopeMap(k.Scopes), ModelACLSet: len(k.AllowedModels) > 0 || len(k.DeniedModels) > 0, ModelAccess: config.ModelAccessRule{Mode: "allow_all", AllowedModels: append([]string(nil), k.AllowedModels...), DeniedModels: append([]string(nil), k.DeniedModels...)}, ServiceClass: strings.TrimSpace(k.ServiceClass)}
if id.Tenant == "" {
return nil, fmt.Errorf("api key %q has no tenant", k.Name)
}
if id.Subject == "" {
id.Subject = "apikey:" + k.Name
}
h := sha256.Sum256([]byte(k.Key))
a.keys[h] = managedKey{identity: id, info: APIKeyInfo{Name: k.Name, Tenant: id.Tenant, Subject: id.Subject, Application: id.Application, Scopes: sortedScopes(k.Scopes), AllowedModels: append([]string(nil), k.AllowedModels...), DeniedModels: append([]string(nil), k.DeniedModels...), ServiceClass: strings.TrimSpace(k.ServiceClass), Source: "config", Deletable: false}}
}
if store != nil {
records, err := store.Load()
if err != nil {
return nil, fmt.Errorf("load persistent API keys: %w", err)
}
for _, rec := range records {
hb, err := hex.DecodeString(rec.HashHex)
if err != nil || len(hb) != sha256.Size {
return nil, fmt.Errorf("persistent API key %q has invalid hash", rec.ID)
}
var h [32]byte
copy(h[:], hb)
if rec.ID == "" || rec.Tenant == "" {
return nil, fmt.Errorf("persistent API key has missing id or tenant")
}
created := rec.CreatedAt
info := APIKeyInfo{ID: rec.ID, Name: rec.Name, Tenant: rec.Tenant, Subject: rec.Subject, Application: rec.Application, Scopes: sortedScopes(rec.Scopes), AllowedModels: append([]string(nil), rec.AllowedModels...), DeniedModels: append([]string(nil), rec.DeniedModels...), ServiceClass: rec.ServiceClass, Source: "persistent", KeyHint: rec.KeyHint, CreatedAt: &created, Deletable: true}
id := Identity{Tenant: rec.Tenant, Subject: rec.Subject, Application: rec.Application, AuthType: "api-key", Scopes: scopeMap(rec.Scopes), ModelACLSet: len(rec.AllowedModels) > 0 || len(rec.DeniedModels) > 0, ModelAccess: config.ModelAccessRule{Mode: "allow_all", AllowedModels: append([]string(nil), rec.AllowedModels...), DeniedModels: append([]string(nil), rec.DeniedModels...)}, ServiceClass: rec.ServiceClass}
if id.Subject == "" {
id.Subject = "apikey:" + rec.Name
info.Subject = id.Subject
}
if _, exists := a.keys[h]; exists {
return nil, fmt.Errorf("persistent API key hash collision for %q", rec.ID)
}
if _, exists := a.runtime[rec.ID]; exists {
return nil, fmt.Errorf("duplicate persistent API key id %q", rec.ID)
}
a.keys[h] = managedKey{identity: id, info: info}
a.runtime[rec.ID] = h
}
}
if cfg.OIDC.Enabled {
v, err := NewOIDCVerifier(ctx, cfg.OIDC)
if err != nil {
return nil, err
}
a.oidc = v
}
return a, nil
}
func scopeMap(in []string) map[string]bool {
m := map[string]bool{}
for _, s := range in {
s = strings.TrimSpace(s)
if s != "" {
m[s] = true
}
}
return m
}
func sortedScopes(in []string) []string {
m := map[string]struct{}{}
for _, s := range in {
if s = strings.TrimSpace(s); s != "" {
m[s] = struct{}{}
}
}
out := make([]string, 0, len(m))
for s := range m {
out = append(out, s)
}
sort.Strings(out)
return out
}
func (a *Authenticator) Authenticate(r *http.Request) (Identity, error) {
ip := a.ClientIP(r)
bypassIP := a.PeerIP(r)
if a.bypassForwarded {
bypassIP = ip
}
for _, rule := range a.bypass {
if containsAny(rule.nets, net.ParseIP(bypassIP)) {
id := rule.identity
// Keep the resolved client IP for observability even though the
// authentication decision defaults to the TCP peer address.
id.ClientIP = ip
return id, nil
}
}
token := ""
if x := strings.TrimSpace(r.Header.Get("X-API-Key")); x != "" {
token = x
}
if token == "" {
h := r.Header.Get("Authorization")
if len(h) > 7 && strings.EqualFold(h[:7], "Bearer ") {
token = strings.TrimSpace(h[7:])
}
}
if token != "" {
sum := sha256.Sum256([]byte(token))
a.mu.RLock()
k, ok := a.keys[sum]
a.mu.RUnlock()
if ok {
id := k.identity
id.ClientIP = ip
return id, nil
}
if a.oidc != nil {
id, err := a.oidc.Verify(r.Context(), token)
if err == nil {
id.ClientIP = ip
return id, nil
}
return Identity{}, fmt.Errorf("invalid bearer token: %w", err)
}
}
return Identity{}, ErrUnauthorized
}
var ErrUnauthorized = fmt.Errorf("authentication required")
// CreateAPIKey creates an API key. With a RuntimeKeyStore configured, only the
// key hash and metadata are persisted. The returned secret is the
// only copy of the plaintext key and must be shown to the administrator once.
func (a *Authenticator) CreateAPIKey(in APIKeyCreate) (APIKeyInfo, string, error) {
if a == nil {
return APIKeyInfo{}, "", errors.New("authenticator unavailable")
}
in.Name = strings.TrimSpace(in.Name)
in.Tenant = strings.TrimSpace(in.Tenant)
in.Subject = strings.TrimSpace(in.Subject)
in.Application = strings.TrimSpace(in.Application)
if in.Name == "" || len(in.Name) > 128 {
return APIKeyInfo{}, "", errors.New("name is required and must be at most 128 characters")
}
if in.Tenant == "" || len(in.Tenant) > 256 {
return APIKeyInfo{}, "", errors.New("tenant is required and must be at most 256 characters")
}
if len(in.Subject) > 256 || len(in.Application) > 256 {
return APIKeyInfo{}, "", errors.New("subject and application must be at most 256 characters")
}
if in.Subject == "" {
in.Subject = "apikey:" + in.Name
}
scopes := sortedScopes(in.Scopes)
for _, s := range scopes {
if len(s) > 128 {
return APIKeyInfo{}, "", errors.New("scope must be at most 128 characters")
}
}
if err := config.ValidateModelAccessRule(config.ModelAccessRule{Mode: "allow_all", AllowedModels: in.AllowedModels, DeniedModels: in.DeniedModels}); err != nil {
return APIKeyInfo{}, "", fmt.Errorf("model ACL: %w", err)
}
secretBytes := make([]byte, 32)
if _, err := rand.Read(secretBytes); err != nil {
return APIKeyInfo{}, "", fmt.Errorf("generate key: %w", err)
}
secret := "ofg_" + base64.RawURLEncoding.EncodeToString(secretBytes)
idBytes := make([]byte, 12)
if _, err := rand.Read(idBytes); err != nil {
return APIKeyInfo{}, "", fmt.Errorf("generate key id: %w", err)
}
id := base64.RawURLEncoding.EncodeToString(idBytes)
h := sha256.Sum256([]byte(secret))
createdAt := time.Now().UTC()
info := APIKeyInfo{ID: id, Name: in.Name, Tenant: in.Tenant, Subject: in.Subject, Application: in.Application, Scopes: scopes, AllowedModels: append([]string(nil), in.AllowedModels...), DeniedModels: append([]string(nil), in.DeniedModels...), ServiceClass: strings.TrimSpace(in.ServiceClass), Source: map[bool]string{true: "persistent", false: "runtime"}[a.store != nil], KeyHint: keyHint(secret), CreatedAt: &createdAt, Deletable: true}
identity := Identity{Tenant: in.Tenant, Subject: in.Subject, Application: in.Application, AuthType: "api-key", Scopes: scopeMap(scopes), ModelACLSet: len(in.AllowedModels) > 0 || len(in.DeniedModels) > 0, ModelAccess: config.ModelAccessRule{Mode: "allow_all", AllowedModels: append([]string(nil), in.AllowedModels...), DeniedModels: append([]string(nil), in.DeniedModels...)}, ServiceClass: strings.TrimSpace(in.ServiceClass)}
a.mu.Lock()
defer a.mu.Unlock()
for _, existing := range a.keys {
if existing.info.Name == in.Name && existing.info.Tenant == in.Tenant {
return APIKeyInfo{}, "", fmt.Errorf("an API key named %q already exists for tenant %q", in.Name, in.Tenant)
}
}
if _, exists := a.keys[h]; exists {
return APIKeyInfo{}, "", errors.New("generated API key collision")
}
if _, exists := a.runtime[id]; exists {
return APIKeyInfo{}, "", errors.New("generated API key id collision")
}
if a.store != nil {
rec := StoredAPIKey{ID: id, Name: in.Name, Tenant: in.Tenant, Subject: in.Subject, Application: in.Application, Scopes: scopes, AllowedModels: append([]string(nil), in.AllowedModels...), DeniedModels: append([]string(nil), in.DeniedModels...), ServiceClass: strings.TrimSpace(in.ServiceClass), HashHex: hex.EncodeToString(h[:]), KeyHint: info.KeyHint, CreatedAt: createdAt}
if err := a.store.Put(rec); err != nil {
return APIKeyInfo{}, "", fmt.Errorf("persist API key: %w", err)
}
}
a.keys[h] = managedKey{identity: identity, info: info}
a.runtime[id] = h
return info, secret, nil
}
func keyHint(secret string) string {
if len(secret) <= 12 {
return secret
}
return secret[:8] + "…" + secret[len(secret)-4:]
}
func (a *Authenticator) APIKeys() []APIKeyInfo {
if a == nil {
return nil
}
a.mu.RLock()
out := make([]APIKeyInfo, 0, len(a.keys))
for _, k := range a.keys {
i := k.info
i.Scopes = append([]string(nil), i.Scopes...)
i.AllowedModels = append([]string(nil), i.AllowedModels...)
i.DeniedModels = append([]string(nil), i.DeniedModels...)
out = append(out, i)
}
a.mu.RUnlock()
sort.Slice(out, func(i, j int) bool {
if out[i].Source != out[j].Source {
return out[i].Source < out[j].Source
}
if out[i].Tenant != out[j].Tenant {
return out[i].Tenant < out[j].Tenant
}
return out[i].Name < out[j].Name
})
return out
}
func (a *Authenticator) DeleteAPIKey(id string) (APIKeyInfo, bool, error) {
if a == nil || id == "" {
return APIKeyInfo{}, false, nil
}
a.mu.Lock()
defer a.mu.Unlock()
h, ok := a.runtime[id]
if !ok {
return APIKeyInfo{}, false, nil
}
k, ok := a.keys[h]
if !ok {
delete(a.runtime, id)
return APIKeyInfo{}, false, nil
}
if a.store != nil {
if err := a.store.Delete(id); err != nil {
return APIKeyInfo{}, true, fmt.Errorf("delete persistent API key: %w", err)
}
}
delete(a.keys, h)
delete(a.runtime, id)
return k.info, true, nil
}
func (a *Authenticator) HasPersistentRuntimeStore() bool { return a != nil && a.store != nil }
func (a *Authenticator) RuntimeStoreHealth(ctx context.Context) error {
if a == nil || a.store == nil {
return nil
}
return a.store.Health(ctx)
}
func (a *Authenticator) PeerIP(r *http.Request) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
host = r.RemoteAddr
}
return strings.TrimSpace(host)
}
func (a *Authenticator) ClientIP(r *http.Request) string {
host := a.PeerIP(r)
peer := net.ParseIP(host)
if peer == nil || !containsAny(a.trusted, peer) {
return host
}
parts := strings.Split(r.Header.Get("X-Forwarded-For"), ",")
chain := make([]net.IP, 0, len(parts)+1)
for _, p := range parts {
if ip := net.ParseIP(strings.TrimSpace(p)); ip != nil {
chain = append(chain, ip)
}
}
chain = append(chain, peer)
for i := len(chain) - 1; i >= 0; i-- {
if !containsAny(a.trusted, chain[i]) {
return chain[i].String()
}
}
if len(chain) > 0 {
return chain[0].String()
}
return strings.TrimSpace(host)
}
func containsAny(nets []*net.IPNet, ip net.IP) bool {
if ip == nil {
return false
}
for _, n := range nets {
if n.Contains(ip) {
return true
}
}
return false
}
func (a *Authenticator) OIDCEnabled() bool { return a != nil && a.oidc != nil }
func (a *Authenticator) OIDCBrowserEndpoints() (BrowserEndpoints, bool) {
if a == nil || a.oidc == nil {
return BrowserEndpoints{}, false
}
return a.oidc.BrowserEndpoints(), true
}
func (a *Authenticator) ExchangeOIDCCode(ctx context.Context, code, redirectURI, clientID, clientSecret, verifier string) (TokenExchange, error) {
if a == nil || a.oidc == nil {
return TokenExchange{}, errors.New("OIDC is disabled")
}
return a.oidc.ExchangeCode(ctx, code, redirectURI, clientID, clientSecret, verifier)
}
func (a *Authenticator) VerifyOIDCToken(ctx context.Context, token string) (Identity, error) {
if a == nil || a.oidc == nil {
return Identity{}, errors.New("OIDC is disabled")
}
return a.oidc.Verify(ctx, token)
}
+103
View File
@@ -0,0 +1,103 @@
package auth
import (
"context"
"net/http/httptest"
"strings"
"testing"
"github.com/example/ollama-fair-gateway/internal/config"
)
func TestRuntimeAPIKeyCreateAuthenticateDelete(t *testing.T) {
a, err := New(context.Background(), config.AuthConfig{APIKeys: []config.APIKeyConfig{{Name: "static", Key: "static-secret", Tenant: "ops", Scopes: []string{"gateway:admin"}}}})
if err != nil {
t.Fatal(err)
}
info, secret, err := a.CreateAPIKey(APIKeyCreate{Name: "openwebui", Tenant: "interactive", Application: "openwebui", Scopes: []string{"models:read", "models:read"}})
if err != nil {
t.Fatal(err)
}
if info.ID == "" || secret == "" || !strings.HasPrefix(secret, "ofg_") || !info.Deletable || info.Source != "runtime" {
t.Fatalf("unexpected created key: %#v secret=%q", info, secret)
}
if len(info.Scopes) != 1 || info.Scopes[0] != "models:read" {
t.Fatalf("unexpected scopes: %#v", info.Scopes)
}
r := httptest.NewRequest("GET", "http://gateway/api/tags", nil)
r.Header.Set("Authorization", "Bearer "+secret)
id, err := a.Authenticate(r)
if err != nil {
t.Fatal(err)
}
if id.Tenant != "interactive" || id.Application != "openwebui" || id.Actor() != "app:openwebui" {
t.Fatalf("unexpected identity: %#v", id)
}
keys := a.APIKeys()
if len(keys) != 2 {
t.Fatalf("keys=%d want 2: %#v", len(keys), keys)
}
for _, k := range keys {
if strings.Contains(k.KeyHint, secret) {
t.Fatalf("key listing leaked secret: %#v", k)
}
}
deleted, ok, err := a.DeleteAPIKey(info.ID)
if err != nil {
t.Fatal(err)
}
if !ok || deleted.Name != "openwebui" {
t.Fatalf("delete failed: ok=%v info=%#v", ok, deleted)
}
if _, err := a.Authenticate(r); err == nil {
t.Fatal("deleted runtime key still authenticates")
}
if _, ok, err := a.DeleteAPIKey(info.ID); err != nil || ok {
t.Fatal("second delete unexpectedly succeeded")
}
}
func TestRuntimeAPIKeyRejectsDuplicateTenantName(t *testing.T) {
a, err := New(context.Background(), config.AuthConfig{})
if err != nil {
t.Fatal(err)
}
if _, _, err := a.CreateAPIKey(APIKeyCreate{Name: "client", Tenant: "team"}); err != nil {
t.Fatal(err)
}
if _, _, err := a.CreateAPIKey(APIKeyCreate{Name: "client", Tenant: "team"}); err == nil {
t.Fatal("expected duplicate error")
}
if _, _, err := a.CreateAPIKey(APIKeyCreate{Name: "client", Tenant: "other"}); err != nil {
t.Fatalf("same name in other tenant should be allowed: %v", err)
}
}
func TestRuntimeAPIKeyCarriesModelACL(t *testing.T) {
a, err := New(context.Background(), config.AuthConfig{})
if err != nil {
t.Fatal(err)
}
info, secret, err := a.CreateAPIKey(APIKeyCreate{Name: "limited", Tenant: "team", AllowedModels: []string{"coding", "qwen3:*"}, DeniedModels: []string{"qwen3:70b*"}})
if err != nil {
t.Fatal(err)
}
if len(info.AllowedModels) != 2 || len(info.DeniedModels) != 1 {
t.Fatalf("info=%#v", info)
}
r := httptest.NewRequest("GET", "http://gateway/api/tags", nil)
r.Header.Set("Authorization", "Bearer "+secret)
id, err := a.Authenticate(r)
if err != nil {
t.Fatal(err)
}
if !id.ModelACLSet {
t.Fatal("expected model ACL on identity")
}
if !config.ModelAccessAllowed(id.ModelAccess, "coding") || config.ModelAccessAllowed(id.ModelAccess, "qwen3:70b-q4") {
t.Fatalf("unexpected ACL: %#v", id.ModelAccess)
}
}
+71
View File
@@ -0,0 +1,71 @@
package auth
import (
"context"
"net/http/httptest"
"testing"
"github.com/example/ollama-fair-gateway/internal/config"
)
func bypassTestConfig(useForwarded bool) config.AuthConfig {
return config.AuthConfig{
IPBypassUseForwardedIP: useForwarded,
TrustedProxies: []string{"10.0.0.0/8"},
IPBypass: []config.IPBypassConfig{{
CIDRs: []string{"127.0.0.1/32"},
Tenant: "local",
Subject: "localhost",
Scopes: []string{"gateway:admin"},
}},
}
}
func TestIPBypassDefaultsToDirectPeerNotForwardedHeader(t *testing.T) {
a, err := New(context.Background(), bypassTestConfig(false))
if err != nil {
t.Fatal(err)
}
r := httptest.NewRequest("GET", "http://gateway/admin", nil)
r.RemoteAddr = "10.2.3.4:12345"
r.Header.Set("X-Forwarded-For", "127.0.0.1")
if got := a.ClientIP(r); got != "127.0.0.1" {
t.Fatalf("resolved client ip changed: got %q", got)
}
if _, err := a.Authenticate(r); err == nil {
t.Fatal("spoofed forwarded loopback unexpectedly satisfied ip_bypass")
}
}
func TestIPBypassForwardedCompatibilityMustBeExplicit(t *testing.T) {
a, err := New(context.Background(), bypassTestConfig(true))
if err != nil {
t.Fatal(err)
}
r := httptest.NewRequest("GET", "http://gateway/admin", nil)
r.RemoteAddr = "10.2.3.4:12345"
r.Header.Set("X-Forwarded-For", "127.0.0.1")
id, err := a.Authenticate(r)
if err != nil {
t.Fatal(err)
}
if !id.IsAdmin() || id.AuthType != "ip-bypass" || id.ClientIP != "127.0.0.1" {
t.Fatalf("unexpected identity: %#v", id)
}
}
func TestIPBypassStillAcceptsDirectPeer(t *testing.T) {
a, err := New(context.Background(), bypassTestConfig(false))
if err != nil {
t.Fatal(err)
}
r := httptest.NewRequest("GET", "http://gateway/admin", nil)
r.RemoteAddr = "127.0.0.1:12345"
id, err := a.Authenticate(r)
if err != nil {
t.Fatal(err)
}
if !id.IsAdmin() || id.ClientIP != "127.0.0.1" {
t.Fatalf("unexpected identity: %#v", id)
}
}
+499
View File
@@ -0,0 +1,499 @@
package auth
import (
"context"
"crypto"
"crypto/ecdsa"
"crypto/ed25519"
"crypto/elliptic"
"crypto/rsa"
"crypto/sha256"
"crypto/sha512"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"math/big"
"net/http"
"net/url"
"strings"
"sync"
"time"
"github.com/example/ollama-fair-gateway/internal/config"
)
type discoveryDoc struct {
Issuer string `json:"issuer"`
JWKSURI string `json:"jwks_uri"`
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
}
type jwksDoc struct {
Keys []json.RawMessage `json:"keys"`
}
type jwkHeader struct {
Kty string `json:"kty"`
Kid string `json:"kid"`
Alg string `json:"alg"`
Use string `json:"use"`
N string `json:"n"`
E string `json:"e"`
Crv string `json:"crv"`
X string `json:"x"`
Y string `json:"y"`
}
type jwtHeader struct {
Alg string `json:"alg"`
Kid string `json:"kid"`
Typ string `json:"typ"`
}
type OIDCVerifier struct {
cfg config.OIDCConfig
issuer string
jwksURI string
authorizationEndpoint string
tokenEndpoint string
client *http.Client
mu sync.RWMutex
keys map[string]crypto.PublicKey
lastRefresh time.Time
allowed map[string]bool
adminGroups map[string]bool
}
func NewOIDCVerifier(ctx context.Context, cfg config.OIDCConfig) (*OIDCVerifier, error) {
issuer := strings.TrimSuffix(cfg.Issuer, "/")
v := &OIDCVerifier{cfg: cfg, issuer: issuer, client: &http.Client{Timeout: 10 * time.Second}, keys: map[string]crypto.PublicKey{}, allowed: map[string]bool{}, adminGroups: map[string]bool{}}
for _, a := range cfg.AllowedAlgorithms {
v.allowed[a] = true
}
for _, g := range cfg.AdminGroups {
v.adminGroups[g] = true
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, issuer+"/.well-known/openid-configuration", nil)
if err != nil {
return nil, err
}
resp, err := v.client.Do(req)
if err != nil {
return nil, fmt.Errorf("oidc discovery: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
return nil, fmt.Errorf("oidc discovery: HTTP %d", resp.StatusCode)
}
var d discoveryDoc
if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&d); err != nil {
return nil, err
}
if strings.TrimSuffix(d.Issuer, "/") != issuer {
return nil, fmt.Errorf("oidc discovery issuer mismatch: got %q", d.Issuer)
}
if d.JWKSURI == "" {
return nil, errors.New("oidc discovery returned no jwks_uri")
}
v.jwksURI = d.JWKSURI
v.authorizationEndpoint = d.AuthorizationEndpoint
v.tokenEndpoint = d.TokenEndpoint
if err := v.refresh(ctx, true); err != nil {
return nil, err
}
return v, nil
}
func (v *OIDCVerifier) Verify(ctx context.Context, token string) (Identity, error) {
parts := strings.Split(token, ".")
if len(parts) != 3 {
return Identity{}, errors.New("token is not a JWT")
}
hb, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return Identity{}, err
}
var h jwtHeader
if err := json.Unmarshal(hb, &h); err != nil {
return Identity{}, err
}
if h.Alg == "none" || !v.allowed[h.Alg] {
return Identity{}, fmt.Errorf("JWT algorithm %q is not allowed", h.Alg)
}
key := v.key(h.Kid)
if key == nil {
if err := v.refresh(ctx, false); err != nil {
return Identity{}, err
}
key = v.key(h.Kid)
}
if key == nil {
return Identity{}, fmt.Errorf("no signing key for kid %q", h.Kid)
}
sig, err := base64.RawURLEncoding.DecodeString(parts[2])
if err != nil {
return Identity{}, err
}
signingInput := []byte(parts[0] + "." + parts[1])
if err := verifySignature(h.Alg, key, signingInput, sig); err != nil {
// Some providers rotate key material while reusing a kid. Refresh once
// (subject to the anti-thundering-herd interval) before rejecting.
_ = v.refresh(ctx, false)
key = v.key(h.Kid)
if key == nil {
return Identity{}, err
}
if err2 := verifySignature(h.Alg, key, signingInput, sig); err2 != nil {
return Identity{}, err2
}
}
pb, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return Identity{}, err
}
dec := json.NewDecoder(strings.NewReader(string(pb)))
dec.UseNumber()
claims := map[string]any{}
if err := dec.Decode(&claims); err != nil {
return Identity{}, err
}
now := time.Now()
skew := v.cfg.ClockSkew.Value()
if iss, _ := claims["iss"].(string); strings.TrimSuffix(iss, "/") != v.issuer {
return Identity{}, errors.New("issuer mismatch")
}
if !audContains(claims["aud"], v.cfg.Audience) {
return Identity{}, errors.New("audience mismatch")
}
if exp, ok := numericTime(claims["exp"]); !ok || now.After(exp.Add(skew)) {
return Identity{}, errors.New("token expired or missing exp")
}
if nbf, ok := numericTime(claims["nbf"]); ok && now.Add(skew).Before(nbf) {
return Identity{}, errors.New("token not valid yet")
}
sub, _ := claims["sub"].(string)
if sub == "" {
return Identity{}, errors.New("token has no sub")
}
tenant := claimString(claims, v.cfg.TenantClaim)
if tenant == "" {
tenant = sub
}
app := claimString(claims, v.cfg.ApplicationClaim)
scopes := map[string]bool{}
for _, s := range claimStrings(claims["scope"]) {
for _, p := range strings.Fields(s) {
scopes[p] = true
}
}
for _, s := range claimStrings(claims["scp"]) {
for _, p := range strings.Fields(s) {
scopes[p] = true
}
}
groups := claimStrings(claimValue(claims, v.cfg.GroupsClaim))
for _, g := range groups {
if v.adminGroups[g] {
scopes["gateway:admin"] = true
}
}
return Identity{Tenant: tenant, Subject: sub, Application: app, AuthType: "oidc", Scopes: scopes}, nil
}
func (v *OIDCVerifier) key(kid string) crypto.PublicKey {
v.mu.RLock()
defer v.mu.RUnlock()
if kid != "" {
return v.keys[kid]
}
if len(v.keys) == 1 {
for _, k := range v.keys {
return k
}
}
return nil
}
func (v *OIDCVerifier) refresh(ctx context.Context, force bool) error {
v.mu.Lock()
defer v.mu.Unlock()
if !force && time.Since(v.lastRefresh) < v.cfg.JWKSRefreshMinInterval.Value() {
return nil
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, v.jwksURI, nil)
if err != nil {
return err
}
resp, err := v.client.Do(req)
if err != nil {
return fmt.Errorf("fetch jwks: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
return fmt.Errorf("fetch jwks: HTTP %d", resp.StatusCode)
}
var d jwksDoc
if err := json.NewDecoder(io.LimitReader(resp.Body, 2<<20)).Decode(&d); err != nil {
return err
}
keys := map[string]crypto.PublicKey{}
for _, raw := range d.Keys {
var j jwkHeader
if err := json.Unmarshal(raw, &j); err != nil {
continue
}
k, err := parseJWK(j)
if err != nil {
continue
}
keys[j.Kid] = k
}
if len(keys) == 0 {
return errors.New("jwks contained no supported signing keys")
}
v.keys = keys
v.lastRefresh = time.Now()
return nil
}
func parseJWK(j jwkHeader) (crypto.PublicKey, error) {
dec := base64.RawURLEncoding.DecodeString
switch j.Kty {
case "RSA":
nB, err := dec(j.N)
if err != nil {
return nil, err
}
eB, err := dec(j.E)
if err != nil {
return nil, err
}
e := 0
for _, b := range eB {
e = e<<8 + int(b)
}
if e == 0 {
return nil, errors.New("bad RSA exponent")
}
return &rsa.PublicKey{N: new(big.Int).SetBytes(nB), E: e}, nil
case "EC":
xb, err := dec(j.X)
if err != nil {
return nil, err
}
yb, err := dec(j.Y)
if err != nil {
return nil, err
}
var c elliptic.Curve
switch j.Crv {
case "P-256":
c = elliptic.P256()
case "P-384":
c = elliptic.P384()
case "P-521":
c = elliptic.P521()
default:
return nil, errors.New("unsupported EC curve")
}
x, y := new(big.Int).SetBytes(xb), new(big.Int).SetBytes(yb)
if !c.IsOnCurve(x, y) {
return nil, errors.New("EC key is not on curve")
}
return &ecdsa.PublicKey{Curve: c, X: x, Y: y}, nil
case "OKP":
if j.Crv != "Ed25519" {
return nil, errors.New("unsupported OKP curve")
}
b, err := dec(j.X)
if err != nil {
return nil, err
}
if len(b) != ed25519.PublicKeySize {
return nil, errors.New("bad Ed25519 key")
}
return ed25519.PublicKey(b), nil
default:
return nil, errors.New("unsupported key type")
}
}
func verifySignature(alg string, key crypto.PublicKey, msg, sig []byte) error {
var hash crypto.Hash
var digest []byte
switch alg {
case "RS256", "PS256", "ES256":
h := sha256.Sum256(msg)
hash = crypto.SHA256
digest = h[:]
case "RS384", "PS384", "ES384":
h := sha512.Sum384(msg)
hash = crypto.SHA384
digest = h[:]
case "RS512", "PS512", "ES512":
h := sha512.Sum512(msg)
hash = crypto.SHA512
digest = h[:]
case "EdDSA":
pk, ok := key.(ed25519.PublicKey)
if !ok || !ed25519.Verify(pk, msg, sig) {
return errors.New("invalid JWT signature")
}
return nil
default:
return errors.New("unsupported JWT algorithm")
}
if strings.HasPrefix(alg, "RS") {
pk, ok := key.(*rsa.PublicKey)
if !ok {
return errors.New("JWT key type mismatch")
}
if err := rsa.VerifyPKCS1v15(pk, hash, digest, sig); err != nil {
return errors.New("invalid JWT signature")
}
return nil
}
if strings.HasPrefix(alg, "PS") {
pk, ok := key.(*rsa.PublicKey)
if !ok {
return errors.New("JWT key type mismatch")
}
if err := rsa.VerifyPSS(pk, hash, digest, sig, &rsa.PSSOptions{SaltLength: rsa.PSSSaltLengthEqualsHash, Hash: hash}); err != nil {
return errors.New("invalid JWT signature")
}
return nil
}
pk, ok := key.(*ecdsa.PublicKey)
if !ok {
return errors.New("JWT key type mismatch")
}
n := (pk.Curve.Params().BitSize + 7) / 8
if len(sig) != 2*n {
return errors.New("bad ECDSA JWT signature length")
}
r, s := new(big.Int).SetBytes(sig[:n]), new(big.Int).SetBytes(sig[n:])
if !ecdsa.Verify(pk, digest, r, s) {
return errors.New("invalid JWT signature")
}
return nil
}
func audContains(v any, want string) bool {
for _, s := range claimStrings(v) {
if s == want {
return true
}
}
return false
}
func claimStrings(v any) []string {
switch x := v.(type) {
case string:
return []string{x}
case []any:
out := []string{}
for _, z := range x {
if s, ok := z.(string); ok {
out = append(out, s)
}
}
return out
case []string:
return x
}
return nil
}
func numericTime(v any) (time.Time, bool) {
var n int64
switch x := v.(type) {
case json.Number:
i, e := x.Int64()
if e != nil {
return time.Time{}, false
}
n = i
case float64:
n = int64(x)
case int64:
n = x
default:
return time.Time{}, false
}
return time.Unix(n, 0), true
}
func claimValue(m map[string]any, path string) any {
if path == "" {
return nil
}
var cur any = m
for _, p := range strings.Split(path, ".") {
mm, ok := cur.(map[string]any)
if !ok {
return nil
}
cur = mm[p]
}
return cur
}
func claimString(m map[string]any, path string) string {
if s, ok := claimValue(m, path).(string); ok {
return s
}
return ""
}
type BrowserEndpoints struct {
Authorization string `json:"authorization_endpoint"`
Token string `json:"token_endpoint"`
}
type TokenExchange struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int64 `json:"expires_in"`
IDToken string `json:"id_token,omitempty"`
Scope string `json:"scope,omitempty"`
}
func (v *OIDCVerifier) BrowserEndpoints() BrowserEndpoints {
return BrowserEndpoints{Authorization: v.authorizationEndpoint, Token: v.tokenEndpoint}
}
func (v *OIDCVerifier) ExchangeCode(ctx context.Context, code, redirectURI, clientID, clientSecret, verifier string) (TokenExchange, error) {
if v.tokenEndpoint == "" {
return TokenExchange{}, errors.New("OIDC discovery returned no token_endpoint")
}
form := url.Values{}
form.Set("grant_type", "authorization_code")
form.Set("code", code)
form.Set("redirect_uri", redirectURI)
form.Set("client_id", clientID)
if verifier != "" {
form.Set("code_verifier", verifier)
}
if clientSecret != "" {
form.Set("client_secret", clientSecret)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, v.tokenEndpoint, strings.NewReader(form.Encode()))
if err != nil {
return TokenExchange{}, err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("Accept", "application/json")
resp, err := v.client.Do(req)
if err != nil {
return TokenExchange{}, fmt.Errorf("OIDC token exchange: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
if err != nil {
return TokenExchange{}, err
}
if resp.StatusCode/100 != 2 {
return TokenExchange{}, fmt.Errorf("OIDC token exchange: HTTP %d", resp.StatusCode)
}
var out TokenExchange
if err := json.Unmarshal(body, &out); err != nil {
return TokenExchange{}, err
}
if out.AccessToken == "" {
return TokenExchange{}, errors.New("OIDC token response contains no access_token")
}
return out, nil
}
+63
View File
@@ -0,0 +1,63 @@
package auth
import (
"context"
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"math/big"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/example/ollama-fair-gateway/internal/config"
)
func TestOIDCVerifierRS256(t *testing.T) {
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
var issuer string
mux := http.NewServeMux()
srv := httptest.NewServer(mux)
defer srv.Close()
issuer = srv.URL
mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(map[string]any{"issuer": issuer, "jwks_uri": issuer + "/jwks"})
})
mux.HandleFunc("/jwks", func(w http.ResponseWriter, r *http.Request) {
e := big.NewInt(int64(key.PublicKey.E)).Bytes()
_ = json.NewEncoder(w).Encode(map[string]any{"keys": []any{map[string]any{"kty": "RSA", "kid": "k1", "alg": "RS256", "n": base64.RawURLEncoding.EncodeToString(key.PublicKey.N.Bytes()), "e": base64.RawURLEncoding.EncodeToString(e)}}})
})
v, err := NewOIDCVerifier(context.Background(), config.OIDCConfig{Enabled: true, Issuer: issuer, Audience: "gateway", TenantClaim: "tenant", ApplicationClaim: "azp", GroupsClaim: "groups", AdminGroups: []string{"admins"}, ClockSkew: config.Duration(time.Second), JWKSRefreshMinInterval: config.Duration(time.Second), AllowedAlgorithms: []string{"RS256"}})
if err != nil {
t.Fatal(err)
}
now := time.Now()
tok := signRS256(t, key, "k1", map[string]any{"iss": issuer, "aud": "gateway", "sub": "alice", "tenant": "team-a", "azp": "web", "groups": []string{"admins"}, "exp": now.Add(time.Minute).Unix(), "nbf": now.Add(-time.Second).Unix()})
id, err := v.Verify(context.Background(), tok)
if err != nil {
t.Fatal(err)
}
if id.Tenant != "team-a" || id.Subject != "alice" || id.Application != "web" || !id.IsAdmin() {
t.Fatalf("unexpected identity: %#v", id)
}
}
func signRS256(t *testing.T, key *rsa.PrivateKey, kid string, claims map[string]any) string {
t.Helper()
h, _ := json.Marshal(map[string]any{"alg": "RS256", "typ": "JWT", "kid": kid})
p, _ := json.Marshal(claims)
a := base64.RawURLEncoding.EncodeToString(h) + "." + base64.RawURLEncoding.EncodeToString(p)
sum := sha256.Sum256([]byte(a))
sig, err := rsa.SignPKCS1v15(rand.Reader, key, crypto.SHA256, sum[:])
if err != nil {
t.Fatal(err)
}
return fmt.Sprintf("%s.%s", a, base64.RawURLEncoding.EncodeToString(sig))
}