merged from gramsrv upstream
This commit is contained in:
parent
79c64ee916
commit
21a0856587
651 changed files with 54774 additions and 4590 deletions
|
|
@ -26,127 +26,69 @@ func NewTempAuthKeyBindingStore(db sqlcgen.DBTX) *TempAuthKeyBindingStore {
|
|||
}
|
||||
|
||||
func (s *TempAuthKeyBindingStore) Save(ctx context.Context, b domain.TempAuthKeyBinding) error {
|
||||
if b.ExpiresAt <= 0 || int64(b.ExpiresAt) > math.MaxInt32 {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
return withAuthIdentityTx(ctx, s.db, "save temp auth key binding", func(tx pgx.Tx) error {
|
||||
return s.saveTx(ctx, tx, b)
|
||||
})
|
||||
_, err := s.SaveWithState(ctx, b)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *TempAuthKeyBindingStore) saveTx(ctx context.Context, tx pgx.Tx, b domain.TempAuthKeyBinding) error {
|
||||
rawID := authKeyIDToInt64(b.TempAuthKeyID)
|
||||
permID := b.PermAuthKeyID
|
||||
// Every operation that may bridge temp and permanent rows enters the
|
||||
// permanent identity gate before taking the raw-key row lock. This is the
|
||||
// same gate/order used by selector advance and permanent revocation.
|
||||
if err := lockPermanentAuthIdentities(ctx, tx, []int64{permID}); err != nil {
|
||||
func (s *TempAuthKeyBindingStore) SaveWithState(ctx context.Context, b domain.TempAuthKeyBinding) (domain.TempAuthKeyBindingResult, error) {
|
||||
if b.ExpiresAt <= 0 || int64(b.ExpiresAt) > math.MaxInt32 {
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
var result domain.TempAuthKeyBindingResult
|
||||
err := withAuthIdentityTx(ctx, s.db, "save temp auth key binding", func(tx pgx.Tx) error {
|
||||
var err error
|
||||
result, err = s.saveTx(ctx, tx, b)
|
||||
return err
|
||||
}
|
||||
var (
|
||||
tempExpiry int
|
||||
tempLayer int
|
||||
tempObservationID int64
|
||||
)
|
||||
if err := tx.QueryRow(ctx, `
|
||||
SELECT expires_at, layer, layer_observation_id
|
||||
FROM auth_keys
|
||||
WHERE auth_key_id = $1
|
||||
FOR UPDATE
|
||||
`, rawID).Scan(&tempExpiry, &tempLayer, &tempObservationID); err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
return fmt.Errorf("lock temporary auth key for binding: %w", err)
|
||||
}
|
||||
if tempExpiry <= 0 || tempExpiry != b.ExpiresAt || rawID == permID {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
|
||||
// The raw row serializes first bind and rebind attempts. Read the binding
|
||||
// only after taking that lock, so a concurrent winner is either visible or
|
||||
// still waiting behind us. A different permanent identity is immutable.
|
||||
var currentPermID int64
|
||||
err := tx.QueryRow(ctx, `
|
||||
SELECT perm_auth_key_id
|
||||
FROM temp_auth_key_bindings
|
||||
WHERE temp_auth_key_id = $1
|
||||
`, rawID).Scan(¤tPermID)
|
||||
switch {
|
||||
case err == nil && currentPermID != permID:
|
||||
return store.ErrTempAuthKeyAlreadyBound
|
||||
case err != nil && !errors.Is(err, pgx.ErrNoRows):
|
||||
return fmt.Errorf("read existing temporary auth key binding: %w", err)
|
||||
}
|
||||
|
||||
var (
|
||||
permExpiry int
|
||||
permLayer int
|
||||
permObservationID int64
|
||||
)
|
||||
if err := tx.QueryRow(ctx, `
|
||||
SELECT expires_at, layer, layer_observation_id
|
||||
FROM auth_keys
|
||||
WHERE auth_key_id = $1
|
||||
FOR UPDATE
|
||||
`, permID).Scan(&permExpiry, &permLayer, &permObservationID); err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
return fmt.Errorf("lock permanent auth key for binding: %w", err)
|
||||
}
|
||||
if permExpiry != 0 {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
|
||||
mergedLayer, mergedObservationID, err := store.MergeAuthKeyLayerObservations(
|
||||
tempLayer, tempObservationID,
|
||||
permLayer, permObservationID,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
q := s.q.WithTx(tx)
|
||||
n, err := q.UpsertTempAuthKeyBinding(ctx, sqlcgen.UpsertTempAuthKeyBindingParams{
|
||||
TempAuthKeyID: authKeyIDToInt64(b.TempAuthKeyID),
|
||||
PermAuthKeyID: b.PermAuthKeyID,
|
||||
Nonce: b.Nonce,
|
||||
TempSessionID: b.TempSessionID,
|
||||
ExpiresAt: int32(b.ExpiresAt),
|
||||
EncryptedMessage: b.EncryptedMessage,
|
||||
})
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *TempAuthKeyBindingStore) saveTx(ctx context.Context, tx pgx.Tx, b domain.TempAuthKeyBinding) (domain.TempAuthKeyBindingResult, error) {
|
||||
rawID := authKeyIDToInt64(b.TempAuthKeyID)
|
||||
var (
|
||||
status string
|
||||
mergedLayer int
|
||||
observationID int64
|
||||
)
|
||||
err := tx.QueryRow(ctx, `
|
||||
/* temp_auth_key_bind_atomic */
|
||||
SELECT bind_status, merged_layer, merged_observation_id
|
||||
FROM public.telesrv_bind_temp_auth_key($1, $2, $3, $4, $5, $6)
|
||||
`,
|
||||
rawID,
|
||||
b.PermAuthKeyID,
|
||||
b.Nonce,
|
||||
b.TempSessionID,
|
||||
b.ExpiresAt,
|
||||
b.EncryptedMessage,
|
||||
).Scan(&status, &mergedLayer, &observationID)
|
||||
if err != nil {
|
||||
var pgErr *pgconn.PgError
|
||||
if errors.As(err, &pgErr) && pgErr.Code == "23503" {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrAuthKeyBindingInvalid
|
||||
}
|
||||
return fmt.Errorf("upsert temp auth key binding: %w", err)
|
||||
return domain.TempAuthKeyBindingResult{}, fmt.Errorf("bind temporary auth key atomically: %w", err)
|
||||
}
|
||||
if n == 0 {
|
||||
return store.ErrAuthKeyBindingInvalid
|
||||
switch status {
|
||||
case "ok":
|
||||
if mergedLayer < 0 || observationID < 0 || (observationID > 0 && mergedLayer == 0) {
|
||||
return domain.TempAuthKeyBindingResult{}, fmt.Errorf(
|
||||
"bind temporary auth key atomically: invalid result layer=%d observation=%d",
|
||||
mergedLayer, observationID,
|
||||
)
|
||||
}
|
||||
return domain.TempAuthKeyBindingResult{Layer: mergedLayer, LayerObservationID: observationID}, nil
|
||||
case "already_bound":
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrTempAuthKeyAlreadyBound
|
||||
case "binding_invalid":
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrAuthKeyBindingInvalid
|
||||
case "layer_invalid":
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrAuthKeySessionLayerInvalid
|
||||
case "layer_conflict":
|
||||
return domain.TempAuthKeyBindingResult{}, store.ErrAuthKeySessionLayerConflict
|
||||
default:
|
||||
return domain.TempAuthKeyBindingResult{}, fmt.Errorf("bind temporary auth key atomically: unknown status %q", status)
|
||||
}
|
||||
keyIDs := []int64{rawID, permID}
|
||||
tag, err := tx.Exec(ctx, `
|
||||
UPDATE auth_keys
|
||||
SET layer = $2,
|
||||
layer_observation_id = $3
|
||||
WHERE auth_key_id = ANY($1::bigint[])
|
||||
`, keyIDs, mergedLayer, mergedObservationID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("merge bound auth key layer defaults: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() != int64(len(keyIDs)) {
|
||||
return fmt.Errorf("merge bound auth key layer defaults: updated %d of %d locked keys", tag.RowsAffected(), len(keyIDs))
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
UPDATE authorizations
|
||||
SET layer = $2
|
||||
WHERE auth_key_id = ANY($1::bigint[])
|
||||
`, keyIDs, mergedLayer); err != nil {
|
||||
return fmt.Errorf("mirror bound auth key layer default: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteExpired 实现 store.TempAuthKeyBindingStore:按 auth_keys.expires_at 的部分索引
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue