-
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
Reference in New Issue
Block a user