feat: add database-refreshed identity authorization
This commit is contained in:
1 parent
716a8031b1
commit
c849077591
12 files changed
+1809
-15
No files matched your search
+6
-1
@@ -7,7 +7,8 @@ unchanged until later route-by-route cutover work passes the shared contracts.
|
||||
|
||||
Implemented Modules:
|
||||
|
||||
- `identity`: legacy `zhinian_session` HMAC, parsing, and chunking contract.
|
||||
- `identity`: legacy `zhinian_session` HMAC/chunking plus database-refreshed
|
||||
account, organization, role, and `sessionVersion` authorization.
|
||||
- `postgres`: fail-closed configuration, verified-CA TLS, readiness, and calls
|
||||
to the existing atomic claim and wallet PostgreSQL functions.
|
||||
- `httpapi`: process health and database readiness handlers.
|
||||
@@ -32,3 +33,7 @@ ZHINIAN_DATA_BACKEND=local GO_BACKEND_PORT=8080 ./backend/zhinian-api
|
||||
|
||||
Only `/api/health` and `/api/ready` are implemented in this foundation. No
|
||||
Ingress, Docker, ACK, Secret, or Worker ownership has moved to Go yet.
|
||||
|
||||
The Identity resolver and PostgreSQL authorization-snapshot Adapter are
|
||||
implemented and tested, but no Go login or current-user HTTP route is exposed
|
||||
yet. Route ownership remains with Next.js until a later path-level cutover.
|
||||
@@ -0,0 +1,177 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AuthorizationSnapshotLoader is the Identity module's single persistence
|
||||
// seam. A false found result means that the account does not exist.
|
||||
type AuthorizationSnapshotLoader interface {
|
||||
FindAuthorizationSnapshot(context.Context, string) (AuthorizationSnapshot, bool, error)
|
||||
}
|
||||
|
||||
// AuthorizationSnapshot contains all database-authoritative claims needed to
|
||||
// authorize one signed session.
|
||||
type AuthorizationSnapshot struct {
|
||||
Account AccountSnapshot `json:"account"`
|
||||
Organization *OrganizationSnapshot `json:"organization"`
|
||||
}
|
||||
|
||||
type AccountSnapshot struct {
|
||||
ID string `json:"id"`
|
||||
Phone string `json:"phone"`
|
||||
DisplayName string `json:"displayName"`
|
||||
Role string `json:"role"`
|
||||
OrganizationID string `json:"organizationId,omitempty"`
|
||||
Status string `json:"status"`
|
||||
SessionVersion int `json:"sessionVersion"`
|
||||
}
|
||||
|
||||
type OrganizationSnapshot struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Status string `json:"status"`
|
||||
}
|
||||
|
||||
type RejectionReason string
|
||||
|
||||
const (
|
||||
RejectionInvalidSession RejectionReason = "invalid_session"
|
||||
RejectionClientMismatch RejectionReason = "client_mismatch"
|
||||
RejectionAccountNotFound RejectionReason = "account_not_found"
|
||||
RejectionAccountDisabled RejectionReason = "account_disabled"
|
||||
RejectionSessionVersionMismatch RejectionReason = "session_version_mismatch"
|
||||
RejectionOrganizationRequired RejectionReason = "organization_required"
|
||||
RejectionOrganizationNotActive RejectionReason = "organization_not_active"
|
||||
RejectionInvalidRole RejectionReason = "invalid_role"
|
||||
)
|
||||
|
||||
var ErrUnauthenticated = errors.New("unauthenticated")
|
||||
|
||||
// UnauthenticatedError retains a diagnostic reason while allowing callers to
|
||||
// collapse all authorization denials with errors.Is(err, ErrUnauthenticated).
|
||||
type UnauthenticatedError struct {
|
||||
Reason RejectionReason
|
||||
}
|
||||
|
||||
func (err *UnauthenticatedError) Error() string {
|
||||
return fmt.Sprintf("%s: %s", ErrUnauthenticated, err.Reason)
|
||||
}
|
||||
|
||||
func (err *UnauthenticatedError) Unwrap() error {
|
||||
return ErrUnauthenticated
|
||||
}
|
||||
|
||||
type Resolver struct {
|
||||
loader AuthorizationSnapshotLoader
|
||||
secret string
|
||||
requiredClientID string
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func NewResolver(loader AuthorizationSnapshotLoader, secret, requiredClientID string, now func() time.Time) *Resolver {
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
if requiredClientID == "" {
|
||||
requiredClientID = "platform"
|
||||
}
|
||||
return &Resolver{
|
||||
loader: loader,
|
||||
secret: secret,
|
||||
requiredClientID: requiredClientID,
|
||||
now: now,
|
||||
}
|
||||
}
|
||||
|
||||
// Resolve authenticates the signed cookie, reloads its account authorization
|
||||
// state, and returns a session whose authorization claims all come from the
|
||||
// database snapshot.
|
||||
func (resolver *Resolver) Resolve(ctx context.Context, cookieValue string) (Session, error) {
|
||||
if resolver == nil || resolver.loader == nil || resolver.secret == "" {
|
||||
return Session{}, fmt.Errorf("identity resolver is not configured")
|
||||
}
|
||||
session, err := Parse(cookieValue, resolver.secret, resolver.now())
|
||||
if err != nil {
|
||||
return Session{}, reject(RejectionInvalidSession)
|
||||
}
|
||||
if session.User.ClientID != resolver.requiredClientID {
|
||||
return Session{}, reject(RejectionClientMismatch)
|
||||
}
|
||||
|
||||
snapshot, found, err := resolver.loader.FindAuthorizationSnapshot(ctx, session.User.ID)
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
if !found {
|
||||
return Session{}, reject(RejectionAccountNotFound)
|
||||
}
|
||||
account := snapshot.Account
|
||||
if account.Status != "active" {
|
||||
return Session{}, reject(RejectionAccountDisabled)
|
||||
}
|
||||
if session.SessionVersion != nil && *session.SessionVersion != 0 && *session.SessionVersion != account.SessionVersion {
|
||||
return Session{}, reject(RejectionSessionVersionMismatch)
|
||||
}
|
||||
|
||||
authMode, authorities, validRole := roleClaims(account.Role)
|
||||
if !validRole {
|
||||
return Session{}, reject(RejectionInvalidRole)
|
||||
}
|
||||
if account.Role != "super_admin" {
|
||||
if account.OrganizationID == "" {
|
||||
return Session{}, reject(RejectionOrganizationRequired)
|
||||
}
|
||||
if snapshot.Organization == nil || snapshot.Organization.ID != account.OrganizationID || snapshot.Organization.Status != "active" {
|
||||
return Session{}, reject(RejectionOrganizationNotActive)
|
||||
}
|
||||
}
|
||||
|
||||
currentVersion := account.SessionVersion
|
||||
resolved := Session{
|
||||
Version: session.Version,
|
||||
AuthMode: authMode,
|
||||
IssuedAt: session.IssuedAt,
|
||||
ExpiresAt: session.ExpiresAt,
|
||||
SessionVersion: ¤tVersion,
|
||||
AccessToken: session.AccessToken,
|
||||
TokenType: session.TokenType,
|
||||
User: User{
|
||||
ID: account.ID,
|
||||
Subject: account.ID,
|
||||
Username: account.Phone,
|
||||
Phone: account.Phone,
|
||||
DisplayName: account.DisplayName,
|
||||
ClientID: resolver.requiredClientID,
|
||||
OrganizationID: account.OrganizationID,
|
||||
Role: account.Role,
|
||||
Status: account.Status,
|
||||
Authorities: authorities,
|
||||
Scope: []string{},
|
||||
},
|
||||
}
|
||||
if snapshot.Organization != nil && snapshot.Organization.ID == account.OrganizationID {
|
||||
resolved.User.OrganizationName = snapshot.Organization.Name
|
||||
}
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
func roleClaims(role string) (AuthMode, []string, bool) {
|
||||
switch role {
|
||||
case "user":
|
||||
return AuthModeUser, []string{"ROLE_USER"}, true
|
||||
case "organization_admin":
|
||||
return AuthModeAdmin, []string{"ROLE_ORGANIZATION_ADMIN", "ORGANIZATION_ADMIN"}, true
|
||||
case "super_admin":
|
||||
return AuthModeAdmin, []string{"ROLE_SUPER_ADMIN", "SUPER_ADMIN"}, true
|
||||
default:
|
||||
return "", nil, false
|
||||
}
|
||||
}
|
||||
|
||||
func reject(reason RejectionReason) error {
|
||||
return &UnauthenticatedError{Reason: reason}
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type authorizationFixture struct {
|
||||
RequiredClientID string `json:"requiredClientId"`
|
||||
SessionSecret string `json:"sessionSecret"`
|
||||
NowUnix int64 `json:"nowUnix"`
|
||||
Cases []struct {
|
||||
Name string `json:"name"`
|
||||
Session Session `json:"session"`
|
||||
Snapshot *AuthorizationSnapshot `json:"snapshot"`
|
||||
Expected struct {
|
||||
Outcome string `json:"outcome"`
|
||||
Reason RejectionReason `json:"reason"`
|
||||
LoaderCalls int `json:"loaderCalls"`
|
||||
Session Session `json:"session"`
|
||||
} `json:"expected"`
|
||||
} `json:"cases"`
|
||||
}
|
||||
|
||||
type recordingAuthorizationLoader struct {
|
||||
snapshot AuthorizationSnapshot
|
||||
found bool
|
||||
err error
|
||||
ids []string
|
||||
}
|
||||
|
||||
func (loader *recordingAuthorizationLoader) FindAuthorizationSnapshot(_ context.Context, id string) (AuthorizationSnapshot, bool, error) {
|
||||
loader.ids = append(loader.ids, id)
|
||||
return loader.snapshot, loader.found, loader.err
|
||||
}
|
||||
|
||||
func TestResolverDrivesPlatformAuthorizationContract(t *testing.T) {
|
||||
fixture := loadAuthorizationFixture(t)
|
||||
if len(fixture.Cases) != 14 {
|
||||
t.Fatalf("authorization fixture cases = %d, want 14", len(fixture.Cases))
|
||||
}
|
||||
|
||||
for _, testCase := range fixture.Cases {
|
||||
t.Run(testCase.Name, func(t *testing.T) {
|
||||
cookie, err := json.Marshal(testCase.Session)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal session fixture: %v", err)
|
||||
}
|
||||
signed, err := Sign(cookie, fixture.SessionSecret)
|
||||
if err != nil {
|
||||
t.Fatalf("sign session fixture: %v", err)
|
||||
}
|
||||
|
||||
loader := &recordingAuthorizationLoader{}
|
||||
if testCase.Snapshot != nil {
|
||||
loader.snapshot = *testCase.Snapshot
|
||||
loader.found = true
|
||||
}
|
||||
resolver := NewResolver(loader, fixture.SessionSecret, fixture.RequiredClientID, func() time.Time {
|
||||
return time.Unix(fixture.NowUnix, 0)
|
||||
})
|
||||
|
||||
got, resolveErr := resolver.Resolve(context.Background(), signed)
|
||||
if len(loader.ids) != testCase.Expected.LoaderCalls {
|
||||
t.Fatalf("loader calls = %d, want %d", len(loader.ids), testCase.Expected.LoaderCalls)
|
||||
}
|
||||
if len(loader.ids) == 1 && loader.ids[0] != testCase.Session.User.ID {
|
||||
t.Fatalf("loader ID = %q, want %q", loader.ids[0], testCase.Session.User.ID)
|
||||
}
|
||||
|
||||
switch testCase.Expected.Outcome {
|
||||
case "authenticated":
|
||||
if resolveErr != nil {
|
||||
t.Fatalf("Resolve() error = %v", resolveErr)
|
||||
}
|
||||
if !reflect.DeepEqual(got, testCase.Expected.Session) {
|
||||
t.Errorf("resolved session mismatch\n got: %#v\nwant: %#v", got, testCase.Expected.Session)
|
||||
}
|
||||
case "unauthenticated":
|
||||
var rejection *UnauthenticatedError
|
||||
if !errors.As(resolveErr, &rejection) {
|
||||
t.Fatalf("Resolve() error = %v, want typed unauthenticated rejection", resolveErr)
|
||||
}
|
||||
if !errors.Is(resolveErr, ErrUnauthenticated) {
|
||||
t.Errorf("Resolve() must collapse to ErrUnauthenticated")
|
||||
}
|
||||
if rejection.Reason != testCase.Expected.Reason {
|
||||
t.Errorf("rejection reason = %q, want %q", rejection.Reason, testCase.Expected.Reason)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("unsupported fixture outcome %q", testCase.Expected.Outcome)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolverPropagatesLoaderError(t *testing.T) {
|
||||
fixture := loadAuthorizationFixture(t)
|
||||
databaseErr := errors.New("database unavailable")
|
||||
loader := &recordingAuthorizationLoader{err: databaseErr}
|
||||
resolver := NewResolver(loader, fixture.SessionSecret, fixture.RequiredClientID, func() time.Time {
|
||||
return time.Unix(fixture.NowUnix, 0)
|
||||
})
|
||||
signed := signFixtureSession(t, fixture.Cases[0].Session, fixture.SessionSecret)
|
||||
|
||||
_, err := resolver.Resolve(context.Background(), signed)
|
||||
if !errors.Is(err, databaseErr) {
|
||||
t.Fatalf("Resolve() error = %v, want database error", err)
|
||||
}
|
||||
if errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatal("database error must not collapse to unauthenticated")
|
||||
}
|
||||
if !reflect.DeepEqual(loader.ids, []string{fixture.Cases[0].Session.User.ID}) {
|
||||
t.Fatalf("loader IDs = %#v", loader.ids)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolverRejectsInvalidCookieWithoutLoadingAuthorization(t *testing.T) {
|
||||
fixture := loadAuthorizationFixture(t)
|
||||
loader := &recordingAuthorizationLoader{}
|
||||
resolver := NewResolver(loader, fixture.SessionSecret, fixture.RequiredClientID, func() time.Time {
|
||||
return time.Unix(fixture.NowUnix, 0)
|
||||
})
|
||||
|
||||
_, err := resolver.Resolve(context.Background(), "not-a-signed-session")
|
||||
var rejection *UnauthenticatedError
|
||||
if !errors.As(err, &rejection) || rejection.Reason != RejectionInvalidSession {
|
||||
t.Fatalf("Resolve() error = %v, want invalid-session rejection", err)
|
||||
}
|
||||
if len(loader.ids) != 0 {
|
||||
t.Fatalf("loader calls = %d, want 0", len(loader.ids))
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewResolverDefaultsRequiredPlatformClient(t *testing.T) {
|
||||
fixture := loadAuthorizationFixture(t)
|
||||
loader := &recordingAuthorizationLoader{}
|
||||
resolver := NewResolver(loader, fixture.SessionSecret, "", func() time.Time {
|
||||
return time.Unix(fixture.NowUnix, 0)
|
||||
})
|
||||
signed := signFixtureSession(t, fixture.Cases[5].Session, fixture.SessionSecret)
|
||||
|
||||
_, err := resolver.Resolve(context.Background(), signed)
|
||||
var rejection *UnauthenticatedError
|
||||
if !errors.As(err, &rejection) || rejection.Reason != RejectionClientMismatch {
|
||||
t.Fatalf("Resolve() error = %v, want client mismatch with default platform client", err)
|
||||
}
|
||||
if len(loader.ids) != 0 {
|
||||
t.Fatalf("loader calls = %d, want 0", len(loader.ids))
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolverFailsClosedWhenNotConfigured(t *testing.T) {
|
||||
for _, resolver := range []*Resolver{
|
||||
nil,
|
||||
NewResolver(nil, "secret", "platform", nil),
|
||||
NewResolver(&recordingAuthorizationLoader{}, "", "platform", nil),
|
||||
} {
|
||||
_, err := resolver.Resolve(context.Background(), "cookie")
|
||||
if err == nil || errors.Is(err, ErrUnauthenticated) {
|
||||
t.Fatalf("Resolve() error = %v, want configuration failure", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func loadAuthorizationFixture(t *testing.T) authorizationFixture {
|
||||
t.Helper()
|
||||
raw, err := os.ReadFile("../../../contracts/auth/platform-session-authorization-v1.json")
|
||||
if err != nil {
|
||||
t.Fatalf("read authorization fixture: %v", err)
|
||||
}
|
||||
var fixture authorizationFixture
|
||||
if err := json.Unmarshal(raw, &fixture); err != nil {
|
||||
t.Fatalf("decode authorization fixture: %v", err)
|
||||
}
|
||||
return fixture
|
||||
}
|
||||
|
||||
func signFixtureSession(t *testing.T, session Session, secret string) string {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(session)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal fixture session: %v", err)
|
||||
}
|
||||
signed, err := Sign(raw, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("sign fixture session: %v", err)
|
||||
}
|
||||
return signed
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
|
||||
)
|
||||
|
||||
const FindAuthorizationSnapshotSQL = `SELECT
|
||||
u.id,
|
||||
u.phone,
|
||||
u.display_name,
|
||||
u.role,
|
||||
u.organization_id,
|
||||
u.status,
|
||||
u.session_version,
|
||||
o.id,
|
||||
o.name,
|
||||
o.status
|
||||
FROM public.platform_users AS u
|
||||
LEFT JOIN public.platform_organizations AS o ON o.id = u.organization_id
|
||||
WHERE u.id = $1::text`
|
||||
|
||||
// FindAuthorizationSnapshot loads all database-authoritative identity claims in
|
||||
// one query. Account and organization state are deliberately not filtered so
|
||||
// the Identity module can apply one authorization policy to every result.
|
||||
func (db *Database) FindAuthorizationSnapshot(ctx context.Context, identityKey string) (identity.AuthorizationSnapshot, bool, error) {
|
||||
if db.config.Backend != BackendPostgres || db.querier == nil {
|
||||
return identity.AuthorizationSnapshot{}, false, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend)
|
||||
}
|
||||
rows, err := db.querier.Query(ctx, FindAuthorizationSnapshotSQL, identityKey)
|
||||
if err != nil {
|
||||
return identity.AuthorizationSnapshot{}, false, fmt.Errorf("query authorization snapshot: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.AuthorizationSnapshot{}, false, fmt.Errorf("read authorization snapshot: %w", err)
|
||||
}
|
||||
return identity.AuthorizationSnapshot{}, false, nil
|
||||
}
|
||||
|
||||
var snapshot identity.AuthorizationSnapshot
|
||||
var organizationID sql.NullString
|
||||
var joinedOrganizationID sql.NullString
|
||||
var organizationName sql.NullString
|
||||
var organizationStatus sql.NullString
|
||||
if err := rows.Scan(
|
||||
&snapshot.Account.ID,
|
||||
&snapshot.Account.Phone,
|
||||
&snapshot.Account.DisplayName,
|
||||
&snapshot.Account.Role,
|
||||
&organizationID,
|
||||
&snapshot.Account.Status,
|
||||
&snapshot.Account.SessionVersion,
|
||||
&joinedOrganizationID,
|
||||
&organizationName,
|
||||
&organizationStatus,
|
||||
); err != nil {
|
||||
return identity.AuthorizationSnapshot{}, false, fmt.Errorf("scan authorization snapshot: %w", err)
|
||||
}
|
||||
if organizationID.Valid {
|
||||
snapshot.Account.OrganizationID = organizationID.String
|
||||
}
|
||||
if joinedOrganizationID.Valid {
|
||||
snapshot.Organization = &identity.OrganizationSnapshot{
|
||||
ID: joinedOrganizationID.String,
|
||||
Name: organizationName.String,
|
||||
Status: organizationStatus.String,
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return identity.AuthorizationSnapshot{}, false, fmt.Errorf("read authorization snapshot: %w", err)
|
||||
}
|
||||
return snapshot, true, nil
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
|
||||
)
|
||||
|
||||
func TestFindAuthorizationSnapshotLoadsAccountAndOrganizationInOneQuery(t *testing.T) {
|
||||
const wantSQL = `SELECT
|
||||
u.id,
|
||||
u.phone,
|
||||
u.display_name,
|
||||
u.role,
|
||||
u.organization_id,
|
||||
u.status,
|
||||
u.session_version,
|
||||
o.id,
|
||||
o.name,
|
||||
o.status
|
||||
FROM public.platform_users AS u
|
||||
LEFT JOIN public.platform_organizations AS o ON o.id = u.organization_id
|
||||
WHERE u.id = $1::text`
|
||||
rows := &identityRows{rows: [][]any{{
|
||||
"account-1", "13800138000", "Zhang San", "organization_admin", "organization-1", "active", 7,
|
||||
"organization-1", "Acme", "active",
|
||||
}}}
|
||||
querier := &identityQuerier{rows: rows}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
|
||||
|
||||
got, found, err := db.FindAuthorizationSnapshot(context.Background(), "account-1")
|
||||
if err != nil {
|
||||
t.Fatalf("FindAuthorizationSnapshot() error = %v", err)
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("FindAuthorizationSnapshot() found = false, want true")
|
||||
}
|
||||
if querier.sql != wantSQL || !reflect.DeepEqual(querier.args, []any{"account-1"}) {
|
||||
t.Fatalf("query = %q args = %#v, want exact SQL and identity argument", querier.sql, querier.args)
|
||||
}
|
||||
want := identity.AuthorizationSnapshot{
|
||||
Account: identity.AccountSnapshot{
|
||||
ID: "account-1", Phone: "13800138000", DisplayName: "Zhang San", Role: "organization_admin",
|
||||
OrganizationID: "organization-1", Status: "active", SessionVersion: 7,
|
||||
},
|
||||
Organization: &identity.OrganizationSnapshot{ID: "organization-1", Name: "Acme", Status: "active"},
|
||||
}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("FindAuthorizationSnapshot() = %#v, want %#v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindAuthorizationSnapshotLoadsUnboundSuperAdminWithNoOrganization(t *testing.T) {
|
||||
querier := &identityQuerier{rows: &identityRows{rows: [][]any{{
|
||||
"super-1", "13900139000", "Root", "super_admin", nil, "active", 4,
|
||||
nil, nil, nil,
|
||||
}}}}
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
|
||||
|
||||
got, found, err := db.FindAuthorizationSnapshot(context.Background(), "super-1")
|
||||
if err != nil {
|
||||
t.Fatalf("FindAuthorizationSnapshot() error = %v", err)
|
||||
}
|
||||
if !found || got.Account.OrganizationID != "" || got.Organization != nil {
|
||||
t.Fatalf("FindAuthorizationSnapshot() = (%#v, %v), want unbound account and nil organization", got, found)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindAuthorizationSnapshotReturnsNotFoundForNoAccount(t *testing.T) {
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, &identityQuerier{rows: &identityRows{}})
|
||||
|
||||
got, found, err := db.FindAuthorizationSnapshot(context.Background(), "missing")
|
||||
if err != nil || found || !reflect.DeepEqual(got, identity.AuthorizationSnapshot{}) {
|
||||
t.Fatalf("FindAuthorizationSnapshot() = (%#v, %v, %v), want zero, false, nil", got, found, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindAuthorizationSnapshotPropagatesDatabaseFailures(t *testing.T) {
|
||||
queryErr := errors.New("query failed")
|
||||
scanErr := errors.New("scan failed")
|
||||
rowsErr := errors.New("rows failed")
|
||||
tests := []struct {
|
||||
name string
|
||||
querier *identityQuerier
|
||||
wantErr error
|
||||
}{
|
||||
{name: "query", querier: &identityQuerier{err: queryErr}, wantErr: queryErr},
|
||||
{name: "scan", querier: &identityQuerier{rows: &identityRows{rows: [][]any{{nil}}, scanErr: scanErr}}, wantErr: scanErr},
|
||||
{name: "rows before first row", querier: &identityQuerier{rows: &identityRows{err: rowsErr}}, wantErr: rowsErr},
|
||||
{name: "rows after scan", querier: &identityQuerier{rows: &identityRows{
|
||||
rows: [][]any{{"account-1", "13800138000", "Name", "user", "org-1", "active", 1, "org-1", "Org", "active"}},
|
||||
err: rowsErr,
|
||||
}}, wantErr: rowsErr},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
db := NewDatabase(Config{Backend: BackendPostgres}, test.querier)
|
||||
got, found, err := db.FindAuthorizationSnapshot(context.Background(), "account-1")
|
||||
if !errors.Is(err, test.wantErr) {
|
||||
t.Fatalf("FindAuthorizationSnapshot() error = %v, want wrapping %v", err, test.wantErr)
|
||||
}
|
||||
if found || !reflect.DeepEqual(got, identity.AuthorizationSnapshot{}) {
|
||||
t.Fatalf("FindAuthorizationSnapshot() = (%#v, %v), want fail-closed zero result", got, found)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindAuthorizationSnapshotFailsClosedWithoutPostgres(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
config Config
|
||||
querier Querier
|
||||
}{
|
||||
{name: "local backend", config: Config{Backend: BackendLocal}, querier: &identityQuerier{err: errors.New("must not query")}},
|
||||
{name: "unavailable pool", config: Config{Backend: BackendPostgres}},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
db := NewDatabase(test.config, test.querier)
|
||||
got, found, err := db.FindAuthorizationSnapshot(context.Background(), "account-1")
|
||||
if err == nil || found || !reflect.DeepEqual(got, identity.AuthorizationSnapshot{}) {
|
||||
t.Fatalf("FindAuthorizationSnapshot() = (%#v, %v, %v), want fail-closed error", got, found, err)
|
||||
}
|
||||
if querier, ok := test.querier.(*identityQuerier); ok && querier.called {
|
||||
t.Fatal("FindAuthorizationSnapshot() queried while PostgreSQL was unavailable")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var _ identity.AuthorizationSnapshotLoader = (*Database)(nil)
|
||||
|
||||
type identityQuerier struct {
|
||||
rows *identityRows
|
||||
err error
|
||||
sql string
|
||||
args []any
|
||||
called bool
|
||||
}
|
||||
|
||||
func (q *identityQuerier) Query(_ context.Context, query string, args ...any) (Rows, error) {
|
||||
q.called = true
|
||||
q.sql = query
|
||||
q.args = args
|
||||
return q.rows, q.err
|
||||
}
|
||||
|
||||
type identityRows struct {
|
||||
rows [][]any
|
||||
idx int
|
||||
err error
|
||||
scanErr error
|
||||
}
|
||||
|
||||
func (r *identityRows) Close() {}
|
||||
func (r *identityRows) Err() error { return r.err }
|
||||
func (r *identityRows) Next() bool { return r.idx < len(r.rows) }
|
||||
|
||||
func (r *identityRows) Scan(dest ...any) error {
|
||||
if r.scanErr != nil {
|
||||
return r.scanErr
|
||||
}
|
||||
if r.idx >= len(r.rows) {
|
||||
return errors.New("scan past end")
|
||||
}
|
||||
row := r.rows[r.idx]
|
||||
r.idx++
|
||||
if len(dest) != len(row) {
|
||||
return errors.New("scan arity mismatch")
|
||||
}
|
||||
for index, target := range dest {
|
||||
value := row[index]
|
||||
switch target := target.(type) {
|
||||
case *string:
|
||||
text, ok := value.(string)
|
||||
if !ok {
|
||||
return errors.New("scan string type mismatch")
|
||||
}
|
||||
*target = text
|
||||
case *int:
|
||||
number, ok := value.(int)
|
||||
if !ok {
|
||||
return errors.New("scan int type mismatch")
|
||||
}
|
||||
*target = number
|
||||
case *sql.NullString:
|
||||
if value == nil {
|
||||
*target = sql.NullString{}
|
||||
continue
|
||||
}
|
||||
text, ok := value.(string)
|
||||
if !ok {
|
||||
return errors.New("scan nullable string type mismatch")
|
||||
}
|
||||
*target = sql.NullString{String: text, Valid: true}
|
||||
default:
|
||||
return errors.New("unsupported scan target")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in new issue
Block a user