diff options
Diffstat (limited to 'internal/auth/static.go')
| -rw-r--r-- | internal/auth/static.go | 124 |
1 files changed, 124 insertions, 0 deletions
diff --git a/internal/auth/static.go b/internal/auth/static.go new file mode 100644 index 0000000..495f90c --- /dev/null +++ b/internal/auth/static.go @@ -0,0 +1,124 @@ +package auth + +import ( + "crypto/sha256" + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" + "sync/atomic" + + "aigw/internal/domain" +) + +var ErrUnauthorized = errors.New("invalid or missing API key") + +type Authenticator interface { + Authenticate(*http.Request) (domain.Principal, error) +} + +type KeyRecord struct { + Key string `json:"key"` + KeyID string `json:"key_id"` + TenantID string `json:"tenant_id"` + ProjectID string `json:"project_id"` + Scopes []string `json:"scopes"` +} + +type StaticAuthenticator struct { + state atomic.Pointer[keySnapshot] + allowAnonymous bool +} + +type keySnapshot struct { + keys map[[sha256.Size]byte]domain.Principal +} + +type HashedKeyRecord struct { + Hash [sha256.Size]byte + Principal domain.Principal +} + +func NewStatic(raw string, allowAnonymous bool) (*StaticAuthenticator, error) { + result := &StaticAuthenticator{allowAnonymous: allowAnonymous} + if strings.TrimSpace(raw) == "" { + if allowAnonymous { + result.ReplaceHashed(nil) + return result, nil + } + return nil, errors.New("client API key environment variable is empty") + } + + var records []KeyRecord + if err := json.Unmarshal([]byte(raw), &records); err != nil { + return nil, fmt.Errorf("parse client API keys JSON: %w", err) + } + hashed := make([]HashedKeyRecord, 0, len(records)) + seen := make(map[[sha256.Size]byte]struct{}, len(records)) + for i, record := range records { + if record.Key == "" || record.KeyID == "" || record.TenantID == "" || record.ProjectID == "" { + return nil, fmt.Errorf("client API key record %d requires key, key_id, tenant_id, and project_id", i) + } + hash := sha256.Sum256([]byte(record.Key)) + if _, exists := seen[hash]; exists { + return nil, fmt.Errorf("duplicate client API key at record %d", i) + } + seen[hash] = struct{}{} + hashed = append(hashed, HashedKeyRecord{Hash: hash, Principal: domain.Principal{ + KeyID: record.KeyID, TenantID: record.TenantID, ProjectID: record.ProjectID, + Scopes: append([]string(nil), record.Scopes...), + }}) + } + if len(hashed) == 0 && !allowAnonymous { + return nil, errors.New("at least one client API key is required") + } + result.ReplaceHashed(hashed) + return result, nil +} + +func NewDynamic(records []HashedKeyRecord, allowAnonymous bool) *StaticAuthenticator { + result := &StaticAuthenticator{allowAnonymous: allowAnonymous} + result.ReplaceHashed(records) + return result +} + +func (a *StaticAuthenticator) ReplaceHashed(records []HashedKeyRecord) { + keys := make(map[[sha256.Size]byte]domain.Principal, len(records)) + for _, record := range records { + principal := record.Principal + principal.Scopes = append([]string(nil), principal.Scopes...) + keys[record.Hash] = principal + } + a.state.Store(&keySnapshot{keys: keys}) +} + +func (a *StaticAuthenticator) Authenticate(r *http.Request) (domain.Principal, error) { + key := bearerToken(r.Header.Get("Authorization")) + if key == "" { + key = strings.TrimSpace(r.Header.Get("x-api-key")) + } + if key == "" && a.allowAnonymous { + return domain.Principal{KeyID: "anonymous", TenantID: "anonymous", ProjectID: "anonymous"}, nil + } + if key == "" { + return domain.Principal{}, ErrUnauthorized + } + snapshot := a.state.Load() + if snapshot == nil { + return domain.Principal{}, ErrUnauthorized + } + principal, ok := snapshot.keys[sha256.Sum256([]byte(key))] + if !ok { + return domain.Principal{}, ErrUnauthorized + } + return principal, nil +} + +func bearerToken(header string) string { + parts := strings.Fields(header) + if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { + return "" + } + return parts[1] +} |
