feat: add shared backend authorization seams
This commit is contained in:
1 parent
d0207fcebe
commit
48dd5d07c8
9 files changed
+1086
No files matched your search
@@ -0,0 +1,162 @@
|
||||
// Package publicapi authenticates public API clients and internal workers.
|
||||
// Runtime configuration is parsed once and injected into an Authenticator;
|
||||
// this package deliberately does not read process environment variables.
|
||||
package publicapi
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const ownerPartLimit = 96
|
||||
|
||||
type PublicClient struct {
|
||||
ID string
|
||||
Key string
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
APIKeys string
|
||||
InternalWorkerToken string
|
||||
Production bool
|
||||
}
|
||||
|
||||
type AuthError struct {
|
||||
Status int
|
||||
Message string
|
||||
}
|
||||
|
||||
func (e *AuthError) Error() string {
|
||||
return e.Message
|
||||
}
|
||||
|
||||
type Authenticator struct {
|
||||
clients []PublicClient
|
||||
internalWorkerToken string
|
||||
production bool
|
||||
}
|
||||
|
||||
func NewAuthenticator(config Config) *Authenticator {
|
||||
return &Authenticator{
|
||||
clients: ParseClients(config.APIKeys),
|
||||
internalWorkerToken: strings.TrimSpace(config.InternalWorkerToken),
|
||||
production: config.Production,
|
||||
}
|
||||
}
|
||||
|
||||
func ParseClients(configured string) []PublicClient {
|
||||
entries := strings.FieldsFunc(configured, func(character rune) bool {
|
||||
return character == ',' || character == '\n'
|
||||
})
|
||||
clients := make([]PublicClient, 0, len(entries))
|
||||
for _, rawEntry := range entries {
|
||||
entry := strings.TrimSpace(rawEntry)
|
||||
if entry == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
id := "default"
|
||||
key := entry
|
||||
if separator := strings.IndexByte(entry, ':'); separator >= 0 {
|
||||
id = strings.TrimSpace(entry[:separator])
|
||||
key = strings.TrimSpace(entry[separator+1:])
|
||||
}
|
||||
if id == "" || key == "" {
|
||||
continue
|
||||
}
|
||||
clients = append(clients, PublicClient{ID: id, Key: key})
|
||||
}
|
||||
return clients
|
||||
}
|
||||
|
||||
func (a *Authenticator) Authenticate(request *http.Request) (PublicClient, string, error) {
|
||||
presented := publicCredential(request)
|
||||
if presented == "" {
|
||||
return PublicClient{}, "", &AuthError{Status: http.StatusUnauthorized, Message: "Missing API key."}
|
||||
}
|
||||
for _, client := range a.clients {
|
||||
if secureEqual(client.Key, presented) {
|
||||
return client, OwnerID(client.ID), nil
|
||||
}
|
||||
}
|
||||
return PublicClient{}, "", &AuthError{Status: http.StatusUnauthorized, Message: "Invalid API key."}
|
||||
}
|
||||
|
||||
func (a *Authenticator) AssertInternalWorker(request *http.Request) error {
|
||||
if a.internalWorkerToken == "" && !a.production {
|
||||
return nil
|
||||
}
|
||||
if a.internalWorkerToken == "" {
|
||||
return &AuthError{Status: http.StatusInternalServerError, Message: "Worker token is not configured."}
|
||||
}
|
||||
|
||||
presented := request.Header.Get("x-zhinian-worker-token")
|
||||
if presented == "" {
|
||||
presented = bearerToken(request)
|
||||
}
|
||||
if presented == "" || !secureEqual(a.internalWorkerToken, presented) {
|
||||
return &AuthError{Status: http.StatusUnauthorized, Message: "Invalid worker token."}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func OwnerID(id string) string {
|
||||
part := sanitizeOwnerPart(id)
|
||||
if part == "" {
|
||||
part = "unknown"
|
||||
}
|
||||
return "api:" + part
|
||||
}
|
||||
|
||||
func publicCredential(request *http.Request) string {
|
||||
if token := bearerToken(request); token != "" {
|
||||
return token
|
||||
}
|
||||
return request.Header.Get("x-zhinian-api-key")
|
||||
}
|
||||
|
||||
func bearerToken(request *http.Request) string {
|
||||
authorization := request.Header.Get("authorization")
|
||||
separator := strings.IndexAny(authorization, " \t\r\n\v\f")
|
||||
if separator <= 0 || !strings.EqualFold(authorization[:separator], "Bearer") {
|
||||
return ""
|
||||
}
|
||||
if strings.TrimLeft(authorization[separator:], " \t\r\n\v\f") == authorization[separator:] {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(authorization[separator:])
|
||||
}
|
||||
|
||||
func secureEqual(expected, presented string) bool {
|
||||
if len(expected) != len(presented) {
|
||||
return false
|
||||
}
|
||||
return subtle.ConstantTimeCompare([]byte(expected), []byte(presented)) == 1
|
||||
}
|
||||
|
||||
func sanitizeOwnerPart(value string) string {
|
||||
part := make([]byte, 0, min(len(value), ownerPartLimit))
|
||||
invalidRun := false
|
||||
for _, character := range value {
|
||||
if isOwnerCharacter(character) {
|
||||
invalidRun = false
|
||||
if len(part) < ownerPartLimit {
|
||||
part = append(part, byte(character))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !invalidRun && len(part) < ownerPartLimit {
|
||||
part = append(part, '_')
|
||||
}
|
||||
invalidRun = true
|
||||
}
|
||||
return string(part)
|
||||
}
|
||||
|
||||
func isOwnerCharacter(character rune) bool {
|
||||
return character >= 'A' && character <= 'Z' ||
|
||||
character >= 'a' && character <= 'z' ||
|
||||
character >= '0' && character <= '9' ||
|
||||
strings.ContainsRune("_.:@-", character)
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
package publicapi_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/publicapi"
|
||||
)
|
||||
|
||||
type authErrorFixture struct {
|
||||
Status int `json:"status"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type authFixture struct {
|
||||
Version int `json:"version"`
|
||||
APIKeys string `json:"apiKeys"`
|
||||
Clients []struct {
|
||||
ID string `json:"id"`
|
||||
Key string `json:"key"`
|
||||
} `json:"clients"`
|
||||
AuthenticationCases []struct {
|
||||
Name string `json:"name"`
|
||||
Headers map[string]string `json:"headers"`
|
||||
Expected *struct {
|
||||
Client struct {
|
||||
ID string `json:"id"`
|
||||
Key string `json:"key"`
|
||||
} `json:"client"`
|
||||
Owner string `json:"owner"`
|
||||
} `json:"expected"`
|
||||
Error *authErrorFixture `json:"error"`
|
||||
} `json:"authenticationCases"`
|
||||
OwnerCases []struct {
|
||||
ID string `json:"id"`
|
||||
Owner string `json:"owner"`
|
||||
MaxPartLength int `json:"maxPartLength"`
|
||||
} `json:"ownerCases"`
|
||||
WorkerCases []struct {
|
||||
Name string `json:"name"`
|
||||
Production bool `json:"production"`
|
||||
ConfiguredToken string `json:"configuredToken"`
|
||||
Headers map[string]string `json:"headers"`
|
||||
Allowed bool `json:"allowed"`
|
||||
Error *authErrorFixture `json:"error"`
|
||||
} `json:"workerCases"`
|
||||
}
|
||||
|
||||
func TestParseClientsMatchesSharedContract(t *testing.T) {
|
||||
fixture := loadAuthFixture(t)
|
||||
clients := publicapi.ParseClients(fixture.APIKeys)
|
||||
if len(clients) != len(fixture.Clients) {
|
||||
t.Fatalf("len(ParseClients()) = %d, want %d", len(clients), len(fixture.Clients))
|
||||
}
|
||||
for index, want := range fixture.Clients {
|
||||
if clients[index] != (publicapi.PublicClient{ID: want.ID, Key: want.Key}) {
|
||||
t.Fatalf("client[%d] = %+v, want %+v", index, clients[index], want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticateMatchesSharedContract(t *testing.T) {
|
||||
fixture := loadAuthFixture(t)
|
||||
authenticator := publicapi.NewAuthenticator(publicapi.Config{APIKeys: fixture.APIKeys})
|
||||
for _, testCase := range fixture.AuthenticationCases {
|
||||
t.Run(testCase.Name, func(t *testing.T) {
|
||||
request := httptest.NewRequest("GET", "/api/v1/generations", nil)
|
||||
for name, value := range testCase.Headers {
|
||||
request.Header.Set(name, value)
|
||||
}
|
||||
client, owner, err := authenticator.Authenticate(request)
|
||||
if testCase.Error != nil {
|
||||
assertAuthError(t, err, *testCase.Error)
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate() error = %v", err)
|
||||
}
|
||||
wantClient := publicapi.PublicClient{ID: testCase.Expected.Client.ID, Key: testCase.Expected.Client.Key}
|
||||
if client != wantClient || owner != testCase.Expected.Owner {
|
||||
t.Fatalf("Authenticate() = (%+v, %q), want (%+v, %q)", client, owner, wantClient, testCase.Expected.Owner)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOwnerIDMatchesSharedContract(t *testing.T) {
|
||||
fixture := loadAuthFixture(t)
|
||||
for _, testCase := range fixture.OwnerCases {
|
||||
t.Run(testCase.ID, func(t *testing.T) {
|
||||
owner := publicapi.OwnerID(testCase.ID)
|
||||
if owner != testCase.Owner {
|
||||
t.Fatalf("OwnerID(%q) = %q, want %q", testCase.ID, owner, testCase.Owner)
|
||||
}
|
||||
if testCase.MaxPartLength > 0 && len(owner)-len("api:") != testCase.MaxPartLength {
|
||||
t.Fatalf("owner part length = %d, want %d", len(owner)-len("api:"), testCase.MaxPartLength)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAssertInternalWorkerMatchesSharedContract(t *testing.T) {
|
||||
fixture := loadAuthFixture(t)
|
||||
for _, testCase := range fixture.WorkerCases {
|
||||
t.Run(testCase.Name, func(t *testing.T) {
|
||||
authenticator := publicapi.NewAuthenticator(publicapi.Config{
|
||||
InternalWorkerToken: testCase.ConfiguredToken,
|
||||
Production: testCase.Production,
|
||||
})
|
||||
request := httptest.NewRequest("POST", "/api/internal/worker/tick", nil)
|
||||
for name, value := range testCase.Headers {
|
||||
request.Header.Set(name, value)
|
||||
}
|
||||
err := authenticator.AssertInternalWorker(request)
|
||||
if testCase.Error != nil {
|
||||
assertAuthError(t, err, *testCase.Error)
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("AssertInternalWorker() error = %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticatorUsesOnlyInjectedConfig(t *testing.T) {
|
||||
t.Setenv("ZHINIAN_API_KEYS", "environment:must-not-be-read")
|
||||
t.Setenv("ZHINIAN_INTERNAL_WORKER_TOKEN", "environment-worker-token")
|
||||
t.Setenv("NODE_ENV", "production")
|
||||
|
||||
authenticator := publicapi.NewAuthenticator(publicapi.Config{})
|
||||
request := httptest.NewRequest("GET", "/", nil)
|
||||
request.Header.Set("authorization", "Bearer must-not-be-read")
|
||||
_, _, err := authenticator.Authenticate(request)
|
||||
assertAuthError(t, err, authErrorFixture{Status: 401, Message: "Invalid API key."})
|
||||
if err := authenticator.AssertInternalWorker(httptest.NewRequest("POST", "/", nil)); err != nil {
|
||||
t.Fatalf("development bypass with injected zero config error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertAuthError(t *testing.T, err error, want authErrorFixture) {
|
||||
t.Helper()
|
||||
if err == nil {
|
||||
t.Fatal("error = nil, want typed authentication error")
|
||||
}
|
||||
authError, ok := err.(*publicapi.AuthError)
|
||||
if !ok {
|
||||
t.Fatalf("error type = %T, want *publicapi.AuthError", err)
|
||||
}
|
||||
if authError.Status != want.Status || authError.Message != want.Message || authError.Error() != want.Message {
|
||||
t.Fatalf("error = %+v, want status=%d message=%q", authError, want.Status, want.Message)
|
||||
}
|
||||
}
|
||||
|
||||
func loadAuthFixture(t *testing.T) authFixture {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile("../../../contracts/auth/public-api-auth-v1.json")
|
||||
if err != nil {
|
||||
t.Fatalf("read public API auth fixture: %v", err)
|
||||
}
|
||||
var fixture authFixture
|
||||
if err := json.Unmarshal(data, &fixture); err != nil {
|
||||
t.Fatalf("decode public API auth fixture: %v", err)
|
||||
}
|
||||
if fixture.Version != 1 {
|
||||
t.Fatalf("fixture version = %d, want 1", fixture.Version)
|
||||
}
|
||||
return fixture
|
||||
}
|
||||
Reference in new issue
Block a user