diff --git a/alembic/versions/0002_notes_widen.py b/alembic/versions/0002_notes_widen.py new file mode 100644 index 0000000..ea923b9 --- /dev/null +++ b/alembic/versions/0002_notes_widen.py @@ -0,0 +1,40 @@ +"""widen secrets.notes from 140 to 255 chars + +Revision ID: 0002_notes_widen +Revises: 0001_initial +Create Date: 2026-10-03 + +Downgrade truncates nothing but re-narrows the column; it fails if any row +holds more than 140 chars. +""" + +from collections.abc import Sequence + +import sqlalchemy as sa + +from alembic import op + +revision: str = "0002_notes_widen" +down_revision: str | None = "0001_initial" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.alter_column( + "secrets", + "notes", + existing_type=sa.String(140), + type_=sa.String(255), + existing_nullable=True, + ) + + +def downgrade() -> None: + op.alter_column( + "secrets", + "notes", + existing_type=sa.String(255), + type_=sa.String(140), + existing_nullable=True, + ) \ No newline at end of file diff --git a/gnexus_creds/api.py b/gnexus_creds/api.py index 79baf15..5081e88 100644 --- a/gnexus_creds/api.py +++ b/gnexus_creds/api.py @@ -28,6 +28,7 @@ ExtensionRead, ImportPayload, Page, + RestorePayload, Scope, SecretCreate, SecretFieldIn, @@ -517,9 +518,16 @@ async def delete_account_data( db: Session = Depends(get_db), actor: Actor = Depends(actor_from_request) ) -> None: - actor.require(Scope.write) - db.query(Secret).filter(Secret.user_id == actor.user.id).delete(synchronize_session=False) - audit(db, actor, action="account_data.deleted") + # hard destructive endpoint: a normal write-scope token must not be able + # to wipe an entire vault in one request — admin scope only (the web UI is + # unaffected: the ui channel bypasses scope checks), plus the sensitive + # rate limiter and a count in the audit trail + actor.require(Scope.admin) + check_actor_rate_limit(db, actor, "account_delete") + deleted = ( + db.query(Secret).filter(Secret.user_id == actor.user.id).delete(synchronize_session=False) + ) + audit(db, actor, action="account_data.deleted", metadata={"deleted": deleted}) db.commit() @@ -565,8 +573,15 @@ @admin_router.post("/restore", tags=["admin"], summary="Restore database from backup") -async def admin_restore(payload: dict) -> dict: - result = restore_backup(payload["filename"]) +async def admin_restore(payload: RestorePayload) -> dict: + # the filename must be one the backup listing knows about — a raw path + # would hand psql an arbitrary file + match = next( + (b for b in list_backups() if b["filename"] == payload.filename), None + ) + if match is None: + raise AppError("backup_not_found", "Backup file not found.", status_code=404) + result = restore_backup(match["path"]) return {"restored": True, **result} diff --git a/gnexus_creds/backup.py b/gnexus_creds/backup.py index 6794ebd..eeb383e 100644 --- a/gnexus_creds/backup.py +++ b/gnexus_creds/backup.py @@ -85,6 +85,10 @@ settings = settings or get_settings() url = settings.database_url path = Path(file_path) + # defense in depth: never feed psql/psql-copy a file outside the backup + # dir, even if a caller skips the listing check + if not path.resolve().is_relative_to(_backup_dir(settings).resolve()): + raise FileNotFoundError(f"Not a backup-dir file: {path}") if not path.exists(): raise FileNotFoundError(f"Backup file not found: {path}") diff --git a/gnexus_creds/config.py b/gnexus_creds/config.py index e55d766..5abd872 100644 --- a/gnexus_creds/config.py +++ b/gnexus_creds/config.py @@ -20,6 +20,12 @@ app_name: str = "gnexus-creds" database_url: str = "postgresql+psycopg://gnexus_creds:gnexus_creds@localhost:5432/gnexus_creds" master_key: str = Field(default="change-me-to-a-32-byte-url-safe-key", min_length=16) + # Per-installation random string; with it the master-key derivation is + # scrypt (stretched), without it the legacy plain sha256 is used. + # Generate once and keep next to the master key: + # python -c "import secrets; print(secrets.token_urlsafe(16))" + # Mandatory in production. Never change it after data exists. + master_key_salt: str = "" session_secret: str = Field(default="change-me", min_length=8) session_cookie_name: str = "gnexus_creds_session" session_ttl_seconds: int = 60 * 60 * 24 * 14 @@ -55,6 +61,13 @@ raise ValueError(f"Unsafe production defaults: {', '.join(sorted(unsafe))}.") if self.database_url.startswith("sqlite"): raise ValueError("Production database_url must use PostgreSQL.") + if not self.master_key_salt: + raise ValueError( + "Production requires master_key_salt (key-stretched derivation; " + "see scripts/migrate-crypto.py before setting it on existing data)." + ) + if "*" in self.cors_origins: + raise ValueError("Production cors_origins must list explicit origins.") for name in ("auth_base_url", "auth_redirect_uri", "mcp_resource_url"): value = getattr(self, name) if not value.startswith("https://"): diff --git a/gnexus_creds/crypto.py b/gnexus_creds/crypto.py index e290544..4feef7c 100644 --- a/gnexus_creds/crypto.py +++ b/gnexus_creds/crypto.py @@ -4,12 +4,13 @@ import hashlib import os import secrets +from functools import lru_cache +from cryptography.exceptions import InvalidTag from cryptography.hazmat.primitives.ciphers.aead import AESGCM ALGORITHM = "AESGCM" - def _b64e(value: bytes) -> str: return base64.urlsafe_b64encode(value).decode().rstrip("=") @@ -19,10 +20,34 @@ return base64.urlsafe_b64decode(value + padding) -def derive_master_key(master_key: str) -> bytes: +def _derive_with_salt(master_key: str, salt: str) -> bytes: + # n=2**15, r=8 needs 128*n*r = 32MB; raise maxmem above OpenSSL's + # built-in default limit, or scrypt raises "memory limit exceeded" + return hashlib.scrypt( + master_key.encode(), salt=salt.encode(), n=2**15, r=8, p=1, dklen=32, maxmem=64 * 2**20 + ) + + +def _derive_legacy(master_key: str) -> bytes: return hashlib.sha256(master_key.encode()).digest() +@lru_cache(maxsize=8) +def _derive_cached(master_key: str, salt: str) -> bytes: + # scrypt costs ~100ms; without the cache every request re-pays it + if salt: + return _derive_with_salt(master_key, salt) + return _derive_legacy(master_key) + + +def derive_master_key(master_key: str, salt: str | None = None) -> bytes: + """Master key from the master string and the installation's salt. + + With a salt this is key-stretched (scrypt); without one it falls back to + the legacy plain sha256 so pre-salt installs keep working.""" + return _derive_cached(master_key, salt or "") + + def new_raw_key() -> bytes: return AESGCM.generate_key(bit_length=256) @@ -39,8 +64,8 @@ return AESGCM(key).decrypt(nonce, ciphertext, aad) -def encrypt_text(key: bytes, key_id: str, value: str) -> dict: - nonce, ciphertext = encrypt_bytes(key, value.encode()) +def encrypt_text(key: bytes, key_id: str, value: str, *, aad: str | None = None) -> dict: + nonce, ciphertext = encrypt_bytes(key, value.encode(), aad=aad.encode() if aad else None) return { "algorithm": ALGORITHM, "key_id": key_id, @@ -49,8 +74,18 @@ } -def decrypt_text(key: bytes, payload: dict) -> str: - return decrypt_bytes(key, _b64d(payload["nonce"]), _b64d(payload["ciphertext"])).decode() +def decrypt_text(key: bytes, payload: dict, *, aad: str | None = None) -> str: + ciphertext = _b64d(payload["ciphertext"]) + nonce = _b64d(payload["nonce"]) + if aad: + try: + return decrypt_bytes(key, nonce, ciphertext, aad=aad.encode()).decode() + except InvalidTag: + # pre-AAD rows (written before the context binding) carry no aad — + # try the unbound form before giving up; anything still failing is + # tampering or a corrupt row and raises InvalidTag below + return decrypt_bytes(key, nonce, ciphertext).decode() + return decrypt_bytes(key, nonce, ciphertext).decode() def token_secret() -> str: @@ -58,4 +93,4 @@ def token_hash(token: str) -> str: - return hashlib.sha256(token.encode()).hexdigest() + return hashlib.sha256(token.encode()).hexdigest() \ No newline at end of file diff --git a/gnexus_creds/models.py b/gnexus_creds/models.py index 2c05dfe..0f89547 100644 --- a/gnexus_creds/models.py +++ b/gnexus_creds/models.py @@ -72,7 +72,7 @@ purpose: Mapped[str | None] = mapped_column(String(255)) category: Mapped[str | None] = mapped_column(String(120), index=True) source: Mapped[str | None] = mapped_column(String(255)) - notes: Mapped[str | None] = mapped_column(String(140)) + notes: Mapped[str | None] = mapped_column(String(255)) status: Mapped[str] = mapped_column(String(24), default="actual", index=True) archived: Mapped[bool] = mapped_column(Boolean, default=False, index=True) allow_ui: Mapped[bool] = mapped_column(Boolean, default=True) diff --git a/gnexus_creds/schemas.py b/gnexus_creds/schemas.py index 9509e5a..2aaf1b3 100644 --- a/gnexus_creds/schemas.py +++ b/gnexus_creds/schemas.py @@ -50,7 +50,7 @@ purpose: str | None = Field(default=None, max_length=255) category: str | None = Field(default=None, max_length=120) source: str | None = Field(default=None, max_length=255) - notes: str | None = Field(default=None, max_length=140) + notes: str | None = Field(default=None, max_length=255) tags: list[str] = Field(default_factory=list) status: SecretStatus = SecretStatus.actual archived: bool = False @@ -61,13 +61,20 @@ @field_validator("tags") @classmethod - def normalize_tags(cls, tags: list[str]) -> list[str]: - result = [] - for tag in tags: - value = tag.strip().lower() - if value and value not in result: - result.append(value[:80]) - return result + def normalize_tags_validator(cls, tags: list[str]) -> list[str]: + return normalize_tags(tags) + + +def normalize_tags(tags: list[str]) -> list[str]: + """Strip, lowercase, dedupe and cap tags — the single normalization + source for both the request schemas and the service layer (an in-service + SecretCreate construct is too fragile).""" + result: list[str] = [] + for tag in tags: + value = tag.strip().lower() + if value and value not in result: + result.append(value[:80]) + return result class SecretUpdate(BaseModel): @@ -75,7 +82,7 @@ purpose: str | None = Field(default=None, max_length=255) category: str | None = Field(default=None, max_length=120) source: str | None = Field(default=None, max_length=255) - notes: str | None = Field(default=None, max_length=140) + notes: str | None = Field(default=None, max_length=255) tags: list[str] | None = None status: SecretStatus | None = None archived: bool | None = None @@ -204,7 +211,13 @@ format: str version: int exported_at: datetime - secrets: list[SecretCreate] + # one request, one vault load: an unlimited list is a cheap DoS and an + # easy accidental double-import of an entire export + secrets: list[SecretCreate] = Field(default_factory=list, max_length=1000) + + +class RestorePayload(BaseModel): + filename: str = Field(min_length=1, max_length=255) class StatsRead(BaseModel): diff --git a/gnexus_creds/services.py b/gnexus_creds/services.py index 51329e5..6cdf88d 100644 --- a/gnexus_creds/services.py +++ b/gnexus_creds/services.py @@ -35,8 +35,21 @@ SecretStatus, SecretUpdate, SecretVersionRead, + normalize_tags, ) +# columns /secrets may sort by; everything else is rejected (422) +_SORTABLE_COLUMNS = { + "title", + "purpose", + "category", + "source", + "status", + "archived", + "created_at", + "updated_at", +} + @dataclass class Actor: @@ -47,6 +60,11 @@ user_agent: str | None = None def require(self, scope: Scope) -> None: + # Documented trust boundary: channel "ui" (a session-cookie request + # from the web UI) bypasses every scope check — the UI manages its own + # data, so scope separation only restricts API tokens ("rest"/"mcp"). + # A stolen web session therefore has full rights; tokens with narrow + # scopes are what you hand to integrations. if self.channel == "ui": return if self.api_token is None or scope.value not in self.api_token.scopes: @@ -54,7 +72,8 @@ def _master_key(settings: Settings | None = None) -> bytes: - return crypto.derive_master_key((settings or get_settings()).master_key) + settings = settings or get_settings() + return crypto.derive_master_key(settings.master_key, settings.master_key_salt) def ensure_user_key(db: Session, user: User) -> UserEncryptionKey: @@ -213,7 +232,14 @@ row.count += 1 -def _store_fields(db: Session, user: User, fields: list[SecretFieldIn]) -> list[dict]: +def _store_fields( + db: Session, user: User, fields: list[SecretFieldIn], *, aad: str | None = None +) -> list[dict]: + """Encrypt fields into the version's stored JSON. + + aad binds each ciphertext to the row identity that carries it + ("::") so a ciphertext transplanted between rows of + the same user fails authentication instead of silently decrypting.""" key_id, key = get_user_key(db, user) stored = [] for index, field in enumerate(sorted(fields, key=lambda item: (item.position, item.name))): @@ -224,7 +250,7 @@ "position": field.position if field.position is not None else index, } payload["value"] = ( - crypto.encrypt_text(key, key_id, field.value) if field.encrypted else field.value + crypto.encrypt_text(key, key_id, field.value, aad=aad) if field.encrypted else field.value ) stored.append(payload) return stored @@ -240,7 +266,12 @@ def _public_fields( - fields: list[dict], *, reveal: bool, db: Session | None = None, user: User | None = None + fields: list[dict], + *, + reveal: bool, + db: Session | None = None, + user: User | None = None, + aad: str | None = None, ): key: bytes | None = None if reveal and db is not None and user is not None: @@ -250,7 +281,7 @@ value = None if reveal: if field.get("encrypted"): - value = crypto.decrypt_text(key, field["value"]) if key else None + value = crypto.decrypt_text(key, field["value"], aad=aad) if key else None else: value = field.get("value") elif not field.get("encrypted"): @@ -278,9 +309,14 @@ def serialize_secret( - secret: Secret, *, reveal: bool = False, db: Session | None = None, user: User | None = None + secret: Secret, + *, + reveal: bool = False, + db: Session | None = None, + user: User | None = None, + version: SecretVersion | None = None, ): - version = _current_version(secret) + version = version or _current_version(secret) cls = SecretReveal if reveal else SecretRead base = { "id": secret.id, @@ -297,7 +333,13 @@ "allow_mcp": secret.allow_mcp, "created_at": secret.created_at, "updated_at": secret.updated_at, - "fields": _public_fields(version.fields, reveal=reveal, db=db, user=user), + "fields": _public_fields( + version.fields, + reveal=reveal, + db=db, + user=user, + aad=_version_aad(secret, version) if reveal else None, + ), } if reveal: base["version_id"] = version.id @@ -305,6 +347,11 @@ return cls(**base) +def _version_aad(secret: Secret, version: SecretVersion) -> str: + """AAD context for this version's field ciphertexts.""" + return f"{secret.user_id}:{secret.id}:{version.id}" + + def create_secret(db: Session, actor: Actor, payload: SecretCreate) -> SecretRead: actor.require(Scope.write) secret = Secret( @@ -324,14 +371,19 @@ db.flush() for tag in payload.tags: db.add(SecretTag(secret_id=secret.id, user_id=actor.user.id, name=tag)) - db.add( - SecretVersion( - secret_id=secret.id, - version_number=1, - fields=_store_fields(db, actor.user, payload.fields), - search_text=_field_search_text(payload.fields), - ) + # the version id is generated here (not by the column default) so the + # ciphertexts can be bound to it as aad before the row exists + version = SecretVersion( + id=uuid.uuid4(), + secret_id=secret.id, + version_number=1, + fields=[], + search_text=_field_search_text(payload.fields), ) + version.fields = _store_fields( + db, actor.user, payload.fields, aad=_version_aad(secret, version) + ) + db.add(version) audit(db, actor, action="secret.created", secret_id=secret.id, metadata={"title": secret.title}) db.flush() db.refresh(secret, attribute_names=["versions", "tags"]) @@ -364,7 +416,8 @@ sort_dir: str = "desc", ) -> tuple[list[SecretRead], int]: actor.require(Scope.read) - limit = min(max(limit, 1), 50) + # api.py's Query allows up to 200; the extension uses the plain 50 + limit = min(max(limit, 1), 200) stmt = select(Secret).where(Secret.user_id == actor.user.id) if not include_archived: stmt = stmt.where(Secret.archived.is_(False)) @@ -387,7 +440,15 @@ Secret.versions.any(SecretVersion.search_text.ilike(like)), ) ) - sort_column = getattr(Secret, sort_by, Secret.updated_at) + # whitelist: getattr(Secret, ...) would also resolve relationships + # ("versions", "tags") and column attributes — both produce a 500 + sort_column = Secret.updated_at + if sort_by in _SORTABLE_COLUMNS: + sort_column = getattr(Secret, sort_by) + elif sort_by: + raise AppError( + "invalid_sort", f"Unknown sort_by value: {sort_by}.", status_code=422 + ) stmt = stmt.order_by(sort_column.desc() if sort_dir == "desc" else sort_column.asc()) count_stmt = select(func.count()).select_from(stmt.subquery()) total = db.scalar(count_stmt) or 0 @@ -432,10 +493,11 @@ secret_id=secret.id, metadata={"title": secret.title, "version_number": version.version_number}, ) - base = serialize_secret(secret, reveal=True, db=db, user=actor.user).model_dump() + base = serialize_secret( + secret, reveal=True, db=db, user=actor.user, version=version + ).model_dump() base["version_id"] = version.id base["version_number"] = version.version_number - base["fields"] = _public_fields(version.fields, reveal=True, db=db, user=actor.user) return SecretReveal(**base) @@ -464,23 +526,28 @@ changed_metadata["status"] = {"old": secret.status, "new": payload.status.value} secret.status = payload.status.value if payload.tags is not None: - normalized = SecretCreate(title="x", tags=payload.tags).tags - if normalized != _tags(secret): + normalized = normalize_tags(payload.tags) + old_tags = _tags(secret) + if normalized != old_tags: + changed_metadata["tags"] = {"old": old_tags, "new": normalized} secret.tags.clear() db.flush() for tag in normalized: db.add(SecretTag(secret_id=secret.id, user_id=actor.user.id, name=tag)) - changed_metadata["tags"] = {"old": _tags(secret), "new": normalized} if payload.fields is not None: - next_version = _current_version(secret).version_number + 1 - db.add( - SecretVersion( - secret_id=secret.id, - version_number=next_version, - fields=_store_fields(db, actor.user, payload.fields), - search_text=_field_search_text(payload.fields), - ) + old_version = _current_version(secret) + next_version = old_version.version_number + 1 + version = SecretVersion( + id=uuid.uuid4(), + secret_id=secret.id, + version_number=next_version, + fields=[], + search_text=_field_search_text(payload.fields), ) + version.fields = _store_fields( + db, actor.user, payload.fields, aad=_version_aad(secret, version) + ) + db.add(version) audit( db, actor, diff --git a/scripts/migrate-crypto.py b/scripts/migrate-crypto.py new file mode 100644 index 0000000..ef17813 --- /dev/null +++ b/scripts/migrate-crypto.py @@ -0,0 +1,130 @@ +"""Crypto migration: scrypt master derivation + AAD-bound field ciphertexts. + +Re-encrypts two layers in the live database: + 1. every wrapped user key (user_encryption_keys) — unwrapped with the + legacy sha256 master and re-wrapped with the salted scrypt derivation; + 2. every encrypted field (secret_versions.fields) — decrypted and + re-encrypted with an aad binding it to "::". + +Idempotent: already-migrated items are detected and skipped, so re-running +after an interruption only finishes the rest. Requires the NEW salt to be +set (GNEXUS_CREDS_MASTER_KEY_SALT); the master key itself must be the same +env value as before the migration. + +Run from the project root with the project venv: + .venv/bin/python scripts/migrate-crypto.py [--dry-run] +""" + +from __future__ import annotations + +import argparse +import sys + +from cryptography.exceptions import InvalidTag + +from gnexus_creds import crypto +from gnexus_creds.crypto import _b64d +from sqlalchemy.orm.attributes import flag_modified +from gnexus_creds.config import get_settings +from gnexus_creds.db import SessionLocal +from gnexus_creds.models import Secret, UserEncryptionKey + + +def _rebind_version_fields(key: bytes, version, aad: str, dry: bool) -> int: + """Return how many ciphertexts were rebound (0 for error rows).""" + rebound = 0 + for field in version.fields: + if not field.get("encrypted"): + continue + envelope = field["value"] + # not crypto.decrypt_text here: its legacy fallback would accept the + # unbound form as "bound" and skip the rebind — probe auth directly + nonce, ciphertext = _b64d(envelope["nonce"]), _b64d(envelope["ciphertext"]) + try: + crypto.decrypt_bytes(key, nonce, ciphertext, aad=aad.encode()) + continue # already bound + except InvalidTag: + pass + try: + plaintext = crypto.decrypt_bytes(key, nonce, ciphertext).decode() + except InvalidTag: + print(" !! cannot decrypt an encrypted field — left as is", file=sys.stderr) + continue + if not dry: + field["value"] = crypto.encrypt_text(key, envelope["key_id"], plaintext, aad=aad) + rebound += 1 + return rebound + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--dry-run", action="store_true", help="report without writing") + args = parser.parse_args() + + settings = get_settings() + if not settings.master_key_salt: + print( + "GNEXUS_CREDS_MASTER_KEY_SALT is not set — nothing to migrate to.", + file=sys.stderr, + ) + return 2 + new_master = crypto.derive_master_key(settings.master_key, settings.master_key_salt) + legacy_master = crypto._derive_legacy(settings.master_key) + + db = SessionLocal() + keys_rewrapped = keys_ok = fields_rebound = 0 + + # 1. user keys: legacy-sha256-wrapped -> scrypt-wrapped + rows = list(db.query(UserEncryptionKey)) + keys = {} + for row in rows: + try: + # already wrapped under the new derivation + keys[row.user_id] = crypto.decrypt_bytes(new_master, row.nonce, row.encrypted_key) + keys_ok += 1 + continue + except InvalidTag: + pass + try: + keys[row.user_id] = crypto.decrypt_bytes(legacy_master, row.nonce, row.encrypted_key) + except InvalidTag: + print( + f"!! cannot unwrap key {row.key_id} (user {row.user_id}) with either master" + " — skipped, its data stays untouched", + file=sys.stderr, + ) + continue + if not args.dry_run: + nonce, encrypted_key = crypto.encrypt_bytes(new_master, keys[row.user_id]) + row.nonce = nonce + row.encrypted_key = encrypted_key + keys_rewrapped += 1 + if not args.dry_run: + db.commit() + + # 2. field ciphertexts: unbound legacy -> aad "" + for secret in db.query(Secret): + key = keys.get(secret.user_id) + if key is None: + continue + for version in secret.versions: + aad = f"{secret.user_id}:{secret.id}:{version.id}" + rebound = _rebind_version_fields(key, version, aad, args.dry_run) + if rebound and not args.dry_run: + # equal-content reassignment is not a change for SQLAlchemy's + # history — force the JSON blob flush properly + flag_modified(version, "fields") + fields_rebound += rebound + if not args.dry_run: + db.commit() + + mode = "DRY RUN: " if args.dry_run else "" + print( + f"{mode}keys rewrapped {keys_rewrapped} (already ok: {keys_ok}); " + f"field ciphertexts rebound {fields_rebound}" + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) \ No newline at end of file diff --git a/tests/test_api.py b/tests/test_api.py index 1f2384c..edb7dc2 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -343,3 +343,92 @@ response = await client.patch("/api/v1/me", json={"locale": "de"}) assert response.status_code == 422 + + + +@pytest.mark.anyio +async def test_account_data_requires_admin_scope(app, auth_app, db_session, user): + write_token = "gcr_write_secret" + db_session.add( + ApiToken( + user_id=user.id, + public_id="w", + name="w", + token_hash=crypto.token_hash(write_token), + scopes=["read", "write"], + ) + ) + db_session.commit() + headers = {"Authorization": f"Bearer {write_token}"} + async with AsyncClient(transport=ASGITransport(app=auth_app), base_url="http://test") as client: + created = await client.post( + "/api/v1/secrets", + json={"title": "Keep", "fields": [{"name": "k", "value": "v", "encrypted": False}]}, + headers=headers, + ) + assert created.status_code == 200 + denied = await client.delete("/api/v1/account-data", headers=headers) + assert denied.status_code == 403 + + # the ui channel manages its own data: the web UI still can wipe everything + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + wiped = await client.delete("/api/v1/account-data") + assert wiped.status_code == 204 + assert (await client.get("/api/v1/secrets")).json()["total"] == 0 + db_session.expire_all() + audit_rows = ( + db_session.query(AuditEvent).filter_by(action="account_data.deleted").all() + ) + assert audit_rows and audit_rows[-1].audit_metadata["deleted"] == 1 + + +@pytest.mark.anyio +async def test_import_caps_secrets_list(app): + item = {"title": "Bulk", "fields": [{"name": "f", "value": "v", "encrypted": False}]} + payload = { + "format": "gnexus-creds-export", + "version": 1, + "exported_at": "2026-10-03T00:00:00Z", + "secrets": [item] * 1001, + } + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/api/v1/import", json=payload) + assert response.status_code == 422 + response = await client.post( + "/api/v1/import", json=dict(payload, secrets=[item] * 1000) + ) + assert response.status_code == 200, response.text + assert response.json() == {"created": 1000} + + +@pytest.mark.anyio +async def test_admin_restore_validates_filename(app): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + missing_field = await client.post("/api/v1/admin/restore", json={}) + assert missing_field.status_code == 422 + unknown = await client.post( + "/api/v1/admin/restore", json={"filename": "backup_unknown.sql"} + ) + assert unknown.status_code == 404 + traversal = await client.post( + "/api/v1/admin/restore", json={"filename": "../backups/backup_unknown.sql"} + ) + assert traversal.status_code == 404 + + +@pytest.mark.anyio +async def test_notes_255_over_api(app): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/api/v1/secrets", json={"title": "N", "notes": "x" * 255}) + assert response.status_code == 200 + too_long = await client.post("/api/v1/secrets", json={"title": "N", "notes": "x" * 256}) + assert too_long.status_code == 422 + + +@pytest.mark.anyio +async def test_unknown_sort_by_is_422(app): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + bad = await client.get("/api/v1/secrets", params={"sort_by": "versions"}) + assert bad.status_code == 422 + ok = await client.get("/api/v1/secrets", params={"sort_by": "title"}) + assert ok.status_code == 200 diff --git a/tests/test_backup.py b/tests/test_backup.py index ebe2dd6..b461331 100644 --- a/tests/test_backup.py +++ b/tests/test_backup.py @@ -136,12 +136,15 @@ class TestRestoreBackupSqlite: def test_copies_backup_to_db_path(self, tmp_path): - backup_file = tmp_path / "backup_20260101_120000.db" + backup_dir = tmp_path / "backups" + backup_dir.mkdir() + backup_file = backup_dir / "backup_20260101_120000.db" backup_file.write_text("restored content") db_path = tmp_path / "target.db" settings = MagicMock() settings.database_url = f"sqlite:///{db_path}" + settings.backup_dir = str(backup_dir) result = restore_backup(str(backup_file), settings) @@ -158,11 +161,14 @@ class TestRestoreBackupPostgresql: def test_runs_psql(self, tmp_path): - backup_file = tmp_path / "backup.sql" + backup_dir = tmp_path / "backups" + backup_dir.mkdir() + backup_file = backup_dir / "backup.sql" backup_file.write_text("-- restore script") settings = MagicMock() settings.database_url = "postgresql+psycopg://user:pass@localhost/db" + settings.backup_dir = str(backup_dir) with patch("gnexus_creds.backup.subprocess.run") as mock_run: mock_run.return_value = MagicMock(returncode=0) @@ -176,11 +182,14 @@ assert str(backup_file) in args def test_raises_on_psql_failure(self, tmp_path): - backup_file = tmp_path / "backup.sql" + backup_dir = tmp_path / "backups" + backup_dir.mkdir() + backup_file = backup_dir / "backup.sql" backup_file.write_text("-- restore script") settings = MagicMock() settings.database_url = "postgresql+psycopg://user:pass@localhost/db" + settings.backup_dir = str(backup_dir) with patch("gnexus_creds.backup.subprocess.run") as mock_run: mock_run.return_value = MagicMock(returncode=1, stderr="restore failed") @@ -198,10 +207,13 @@ create_backup(settings) def test_restore_backup_raises(self, tmp_path): - backup_file = tmp_path / "backup.sql" + backup_dir = tmp_path / "backups" + backup_dir.mkdir() + backup_file = backup_dir / "backup.sql" backup_file.write_text("-- script") settings = MagicMock() settings.database_url = "mysql://localhost/db" + settings.backup_dir = str(backup_dir) with pytest.raises(RuntimeError, match="Unsupported database URL"): restore_backup(str(backup_file), settings) diff --git a/tests/test_config.py b/tests/test_config.py index 771d0a8..e9b180c 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -14,6 +14,8 @@ auth_client_id="gnexus-creds", auth_client_secret="change-me", auth_webhook_secret="change-me", + master_key_salt="prod-salt-value-0", + cors_origins=["https://creds.gnexus.space"], auth_base_url="https://auth.gnexus.space", auth_redirect_uri="https://creds.gnexus.space/auth/callback", mcp_resource_url="https://creds.gnexus.space/mcp-protocol/", @@ -30,6 +32,8 @@ auth_client_id="prod-client", auth_client_secret="prod-client-secret", auth_webhook_secret="prod-webhook-secret", + master_key_salt="prod-salt-value-1", + cors_origins=["https://creds.gnexus.space"], auth_base_url="https://auth.gnexus.space", auth_redirect_uri="https://creds.gnexus.space/auth/callback", mcp_resource_url="https://creds.gnexus.space/mcp-protocol/", @@ -46,7 +50,42 @@ auth_client_id="prod-client", auth_client_secret="prod-client-secret", auth_webhook_secret="prod-webhook-secret", + master_key_salt="prod-salt-value-2", + cors_origins=["https://creds.gnexus.space"], auth_base_url="http://auth.gnexus.space", auth_redirect_uri="https://creds.gnexus.space/auth/callback", mcp_resource_url="https://creds.gnexus.space/mcp-protocol/", ) + + +def test_production_requires_master_key_salt(): + with pytest.raises(ValidationError, match="master_key_salt"): + Settings( + env="production", + database_url="postgresql+psycopg://user:pass@postgres:5432/db", + master_key="prod-master-key-prod-master-key", + session_secret="prod-session-secret", + auth_client_id="prod-client", + auth_client_secret="prod-client-secret", + auth_webhook_secret="prod-webhook-secret", + auth_base_url="https://auth.gnexus.space", + auth_redirect_uri="https://creds.gnexus.space/auth/callback", + mcp_resource_url="https://creds.gnexus.space/mcp-protocol/", + ) + + +def test_production_rejects_wildcard_cors(): + with pytest.raises(ValidationError, match="cors_origins"): + Settings( + env="production", + database_url="postgresql+psycopg://user:pass@postgres:5432/db", + master_key="prod-master-key-prod-master-key", + master_key_salt="prod-salt-value-3", + session_secret="prod-session-secret", + auth_client_id="prod-client", + auth_client_secret="prod-client-secret", + auth_webhook_secret="prod-webhook-secret", + auth_base_url="https://auth.gnexus.space", + auth_redirect_uri="https://creds.gnexus.space/auth/callback", + mcp_resource_url="https://creds.gnexus.space/mcp-protocol/", + ) diff --git a/tests/test_core.py b/tests/test_core.py index 348e366..0f2d87e 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -56,3 +56,68 @@ now = datetime.now(UTC) naive_start = (now - timedelta(minutes=2)).replace(tzinfo=None) assert is_expired(naive_start, now=now, delta=timedelta(minutes=1)) + + +def test_reveal_works_and_transplant_is_caught(db_session, actor): + import pytest as _pytest + from cryptography.exceptions import InvalidTag + + from gnexus_creds.models import SecretVersion, utcnow, Secret as _Secret # noqa: F401 + + created = create_secret( + db_session, + actor, + SecretCreate( + title="Vault", + fields=[SecretFieldIn(name="password", value="v1", encrypted=True, position=1)], + ), + ) + db_session.commit() + update_secret( + db_session, + actor, + created.id, + SecretUpdate(fields=[SecretFieldIn(name="password", value="v2", encrypted=True, position=1)]), + ) + db_session.commit() + + versions = list_versions(db_session, actor, created.id) + revealed = reveal_secret(db_session, actor, created.id) + assert {f.name: f.value for f in revealed.fields}["password"] == "v2" + + # swap the two versions' password ciphertexts — same user, same key, so + # only the aad binding can notice + stored = db_session.query(SecretVersion).filter_by(secret_id=created.id).all() + a, b = stored + a.fields, b.fields = ( + [dict(f, value=other["value"]) for f, other in + ((f, next(x for x in b.fields if x["name"] == f["name"])) for f in a.fields)], + [dict(f, value=other["value"]) for f, other in + ((f, next(x for x in a.fields if x["name"] == f["name"])) for f in b.fields)], + ) + db_session.flush() + with _pytest.raises(InvalidTag): + reveal_secret(db_session, actor, created.id) + + +def test_notes_accept_up_to_255_chars(db_session, actor): + long_notes = "н" * 255 + created = create_secret( + db_session, actor, SecretCreate(title="NoteCap", notes=long_notes) + ) + db_session.commit() + assert created.notes == long_notes + + +def test_tags_update_audit_keeps_old_value(db_session, actor): + from gnexus_creds.models import AuditEvent + + created = create_secret( + db_session, actor, SecretCreate(title="Tagged", tags=["old"]) + ) + db_session.commit() + update_secret(db_session, actor, created.id, SecretUpdate(tags=["new"])) + db_session.commit() + row = db_session.query(AuditEvent).filter_by(action="secret.metadata_updated").one() + diff = row.audit_metadata["diff"]["tags"] + assert diff == {"old": ["old"], "new": ["new"]} diff --git a/tests/test_crypto.py b/tests/test_crypto.py new file mode 100644 index 0000000..1e2ad94 --- /dev/null +++ b/tests/test_crypto.py @@ -0,0 +1,61 @@ +import hashlib + +import pytest +from cryptography.exceptions import InvalidTag + +from gnexus_creds import crypto +from gnexus_creds.errors import AppError + + +def test_derive_without_salt_is_legacy_sha256(): + assert crypto.derive_master_key("key") == hashlib.sha256(b"key").digest() + + +def test_derive_with_salt_is_stretched_and_differs(): + legacy = crypto.derive_master_key("key") + salted = crypto.derive_master_key("key", "install-salt") + assert salted != legacy + assert len(salted) == 32 + # same call again hits the cache and stays deterministic + assert crypto.derive_master_key("key", "install-salt") == salted + + +def test_text_envelope_roundtrip_with_aad(): + key = crypto.new_raw_key() + payload = crypto.encrypt_text(key, "uk_abc", "hint", aad="user:secret:version") + assert crypto.decrypt_text(key, payload, aad="user:secret:version") == "hint" + + +def test_legacy_envelope_still_decrypts_with_aad_context(): + # rows written before the aad binding carry no bound context — the + # tolerant decrypt keeps them open (migration rebinds them) + key = crypto.new_raw_key() + payload = crypto.encrypt_text(key, "uk_abc", "legacy") + assert crypto.decrypt_text(key, payload, aad="user:secret:version") == "legacy" + + +def test_transplanted_envelope_fails_authentication(): + key = crypto.new_raw_key() + p1 = crypto.encrypt_text(key, "uk_abc", "one", aad="user:s1:v1") + p2 = crypto.encrypt_text(key, "uk_abc", "two", aad="user:s2:v2") + # transplant p2's ciphertext into p1's envelope (same user, same key) + p1["ciphertext"] = p2["ciphertext"] + p1["nonce"] = p2["nonce"] + with pytest.raises(InvalidTag): + crypto.decrypt_text(key, p1, aad="user:s1:v1") + + +def test_list_secrets_rejects_unknown_sort_by(db_session, actor): + from gnexus_creds.models import Secret + from gnexus_creds.services import list_secrets + + with pytest.raises(AppError) as exc: + list_secrets(db_session, actor, sort_by="__dict__") + assert exc.value.status_code == 422 + with pytest.raises(AppError) as exc: + list_secrets(db_session, actor, sort_by="versions") # a relationship + assert exc.value.status_code == 422 + # whitelist values keep working + list_secrets(db_session, actor, sort_by="title") + list_secrets(db_session, actor, sort_by="updated_at") + assert list_secrets(db_session, actor, limit=200)[1] == 0 \ No newline at end of file