389 lines
15 KiB
Python
389 lines
15 KiB
Python
import contextlib
|
|
import threading
|
|
import time
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
from imagefind.database import SCHEMA_VERSION, Database, DatabaseTransientError
|
|
from imagefind.security import AuthService, LoginRateLimitError, SecretStore
|
|
|
|
|
|
def test_database_instances_use_independent_postgres_transactions(tmp_path: Path):
|
|
path = tmp_path / "coordinated-postgres"
|
|
first = Database(path)
|
|
first.initialize()
|
|
second = Database(path)
|
|
writer_started = threading.Event()
|
|
release_writer = threading.Event()
|
|
second_finished = threading.Event()
|
|
errors: list[Exception] = []
|
|
|
|
def hold_writer():
|
|
try:
|
|
with first.transaction() as conn:
|
|
conn.execute("INSERT INTO settings(key,value,updated_at) VALUES('first','1','now')")
|
|
writer_started.set()
|
|
release_writer.wait(2)
|
|
except Exception as exc: # pragma: no cover - asserted below
|
|
errors.append(exc)
|
|
|
|
def queued_writer():
|
|
try:
|
|
writer_started.wait(2)
|
|
with second.transaction() as conn:
|
|
conn.execute("INSERT INTO settings(key,value,updated_at) VALUES('second','2','now')")
|
|
second_finished.set()
|
|
except Exception as exc: # pragma: no cover - asserted below
|
|
errors.append(exc)
|
|
|
|
holding = threading.Thread(target=hold_writer)
|
|
waiting = threading.Thread(target=queued_writer)
|
|
holding.start()
|
|
assert writer_started.wait(2)
|
|
waiting.start()
|
|
assert second_finished.wait(1)
|
|
assert first.status()["writer_queue_depth"] == 0
|
|
time.sleep(0.26)
|
|
release_writer.set()
|
|
holding.join(2)
|
|
waiting.join(2)
|
|
|
|
assert not errors
|
|
with first.read() as conn:
|
|
keys = {row[0] for row in conn.execute("SELECT key FROM settings WHERE key IN ('first','second')").fetchall()}
|
|
assert keys == {"first", "second"}
|
|
status = second.status()
|
|
assert status["engine"] == "postgresql"
|
|
assert status["journal_mode"] == "server"
|
|
assert status["writer_queue_depth"] == 0
|
|
assert status["writer_active"] is False
|
|
assert status["pool_max"] >= 2
|
|
|
|
|
|
def test_write_retry_recovers_a_transient_postgres_conflict(tmp_path: Path, monkeypatch):
|
|
db = Database(tmp_path / "retry-postgres")
|
|
db.initialize()
|
|
original_transaction = db.transaction
|
|
attempts = 0
|
|
|
|
@contextlib.contextmanager
|
|
def transient_transaction():
|
|
nonlocal attempts
|
|
attempts += 1
|
|
if attempts == 1:
|
|
raise DatabaseTransientError("transient PostgreSQL serialization failure")
|
|
with original_transaction() as conn:
|
|
yield conn
|
|
|
|
monkeypatch.setattr(db, "transaction", transient_transaction)
|
|
db.write_with_retry(
|
|
lambda conn: conn.execute("INSERT INTO settings(key,value,updated_at) VALUES('recovered','1','now')"),
|
|
timeout_seconds=2,
|
|
)
|
|
|
|
assert db.setting("recovered") == 1
|
|
assert attempts == 2
|
|
assert db.status()["lock_retry_count"] >= 1
|
|
|
|
|
|
def test_database_auth_and_encrypted_secrets(tmp_path: Path):
|
|
db = Database(tmp_path / "auth-postgres")
|
|
db.initialize()
|
|
auth = AuthService(db, session_days=1)
|
|
assert not auth.is_configured()
|
|
with pytest.raises(ValueError):
|
|
auth.setup("short")
|
|
auth.setup("correct horse battery staple")
|
|
token, csrf, _ = auth.login("correct horse battery staple")
|
|
assert auth.session(token)["csrf_token"] == csrf
|
|
assert auth.session(token)["auth_source"] == "local"
|
|
gateway_token, gateway_csrf, _ = auth.login_gateway("nas-admin", "Administrator")
|
|
gateway_session = auth.session(gateway_token)
|
|
assert gateway_session["csrf_token"] == gateway_csrf
|
|
assert gateway_session["auth_source"] == "gateway"
|
|
assert gateway_session["external_user_id"] == "nas-admin"
|
|
with pytest.raises(ValueError):
|
|
auth.login("wrong")
|
|
|
|
token_id, api_token = auth.create_api_token("automation")
|
|
assert auth.verify_api_token(api_token)
|
|
assert auth.verify_api_token(api_token)
|
|
with db.read() as conn:
|
|
assert conn.execute("SELECT last_used_at FROM api_tokens WHERE id=?", (token_id,)).fetchone()[0] is None
|
|
assert auth.flush_api_token_usage() == 1
|
|
with db.read() as conn:
|
|
assert conn.execute("SELECT last_used_at FROM api_tokens WHERE id=?", (token_id,)).fetchone()[0]
|
|
|
|
auth.set_password("a replacement password", replace=True)
|
|
assert auth.login("a replacement password")
|
|
assert auth.session(token) is None
|
|
assert auth.verify_api_token(api_token)
|
|
with pytest.raises(ValueError):
|
|
auth.login("correct horse battery staple")
|
|
|
|
auth.revoke_api_token(token_id)
|
|
assert not auth.verify_api_token(api_token)
|
|
|
|
store = SecretStore(tmp_path / "secret.key")
|
|
ciphertext = store.encrypt_json({"password": "not-plaintext"})
|
|
assert "not-plaintext" not in ciphertext
|
|
assert store.decrypt_json(ciphertext) == {"password": "not-plaintext"}
|
|
|
|
|
|
def test_login_rate_limit_is_scoped_by_client_and_success_clears_failures(tmp_path: Path, monkeypatch):
|
|
db = Database(tmp_path / "rate-limit-postgres")
|
|
db.initialize()
|
|
auth = AuthService(db, session_days=1)
|
|
auth.setup("correct horse battery staple")
|
|
monkeypatch.setattr(auth, "LOGIN_MAX_FAILURES", 2)
|
|
|
|
for _ in range(2):
|
|
with pytest.raises(ValueError, match="密码错误"):
|
|
auth.login("wrong password", client_key="192.0.2.10")
|
|
with pytest.raises(LoginRateLimitError) as blocked:
|
|
auth.login("correct horse battery staple", client_key="192.0.2.10")
|
|
assert blocked.value.retry_after > 0
|
|
|
|
# A separate client remains usable, and a successful login clears that
|
|
# client's partial failure history.
|
|
with pytest.raises(ValueError, match="密码错误"):
|
|
auth.login("wrong password", client_key="192.0.2.11")
|
|
token, _, _ = auth.login("correct horse battery staple", client_key="192.0.2.11")
|
|
assert auth.session(token)
|
|
with pytest.raises(ValueError, match="密码错误"):
|
|
auth.login("wrong password", client_key="192.0.2.11")
|
|
|
|
|
|
def test_api_token_scopes_are_enforced_and_legacy_admin_default_is_preserved(tmp_path: Path):
|
|
db = Database(tmp_path / "token-scopes-postgres")
|
|
db.initialize()
|
|
auth = AuthService(db, session_days=1)
|
|
|
|
_, webdav_token = auth.create_api_token("DAV only", ["webdav"])
|
|
_, media_token = auth.create_api_token("Media only", ["media:read"])
|
|
_, admin_token = auth.create_api_token("Administrator")
|
|
|
|
assert auth.verify_api_token(webdav_token, "webdav")
|
|
assert not auth.verify_api_token(webdav_token, "admin")
|
|
assert not auth.verify_api_token(webdav_token, "media:read")
|
|
assert auth.verify_api_token(media_token, "media:read")
|
|
assert not auth.verify_api_token(media_token, "webdav")
|
|
assert auth.verify_api_token(admin_token, "admin")
|
|
assert auth.verify_api_token(admin_token, "webdav")
|
|
assert auth.verify_api_token(admin_token, "media:read")
|
|
with pytest.raises(ValueError, match="权限范围"):
|
|
auth.create_api_token("invalid", ["unknown"])
|
|
|
|
|
|
def test_postgres_pool_is_bounded_and_reaps_extra_idle_connections(tmp_path: Path):
|
|
db = Database(tmp_path / "bounded-pool-postgres")
|
|
db._pool_min = 1
|
|
db._pool_max = 2
|
|
db._pool_timeout = 0.1
|
|
db._pool_idle_timeout = 1
|
|
db.initialize()
|
|
|
|
with db.read() as first, db.read() as second:
|
|
assert first.execute("SELECT 1").fetchone()[0] == 1
|
|
assert second.execute("SELECT 1").fetchone()[0] == 1
|
|
with pytest.raises(TimeoutError, match="连接超时"):
|
|
with db.read():
|
|
pass
|
|
status = db.status()
|
|
assert status["pool_size"] == 2
|
|
assert status["pool_wait_count"] >= 1
|
|
|
|
time.sleep(1.05)
|
|
with db.read() as conn:
|
|
assert conn.execute("SELECT 1").fetchone()[0] == 1
|
|
assert db.status()["pool_size"] == 1
|
|
activity = db.activity()
|
|
assert isinstance(activity["states"], dict)
|
|
assert activity["waiting"] >= 0
|
|
|
|
|
|
def test_password_validation_and_replacement_guard(tmp_path: Path):
|
|
db = Database(tmp_path / "password-postgres")
|
|
db.initialize()
|
|
auth = AuthService(db, session_days=1)
|
|
|
|
with pytest.raises(ValueError, match="至少"):
|
|
auth.set_password("short")
|
|
with pytest.raises(ValueError, match="超过"):
|
|
auth.set_password("x" * 257)
|
|
|
|
auth.set_password("initial administrator password")
|
|
with pytest.raises(ValueError, match="已经初始化"):
|
|
auth.set_password("second administrator password")
|
|
|
|
|
|
def test_postgres_mvcc_reader_stays_responsive_while_writer_is_open(tmp_path: Path):
|
|
db = Database(tmp_path / "mvcc-postgres")
|
|
db.initialize()
|
|
writer = db.connect()
|
|
try:
|
|
writer.execute("INSERT INTO settings(key,value,updated_at) VALUES('writer','1','now')")
|
|
started = time.monotonic()
|
|
with db.read() as reader:
|
|
assert reader.execute("SELECT count(*) FROM settings WHERE key='writer'").fetchone()[0] == 0
|
|
assert time.monotonic() - started < 0.5
|
|
finally:
|
|
writer.rollback()
|
|
writer.close()
|
|
|
|
|
|
def test_slow_reader_does_not_serialize_other_thread_readers(tmp_path: Path):
|
|
db = Database(tmp_path / "parallel-postgres")
|
|
db.initialize()
|
|
first_entered = threading.Event()
|
|
release_first = threading.Event()
|
|
second_finished = threading.Event()
|
|
errors: list[Exception] = []
|
|
|
|
def hold_reader():
|
|
try:
|
|
with db.read() as reader:
|
|
assert reader.execute("SELECT count(*) FROM settings").fetchone()[0] >= 1
|
|
first_entered.set()
|
|
assert release_first.wait(2)
|
|
except Exception as exc: # pragma: no cover - surfaced below
|
|
errors.append(exc)
|
|
|
|
def use_second_reader():
|
|
try:
|
|
assert first_entered.wait(2)
|
|
with db.read() as reader:
|
|
assert reader.execute("SELECT count(*) FROM settings").fetchone()[0] >= 1
|
|
second_finished.set()
|
|
except Exception as exc: # pragma: no cover - surfaced below
|
|
errors.append(exc)
|
|
|
|
first = threading.Thread(target=hold_reader)
|
|
second = threading.Thread(target=use_second_reader)
|
|
first.start()
|
|
second.start()
|
|
assert second_finished.wait(1), "an unrelated reader waited behind the first reader"
|
|
release_first.set()
|
|
first.join(timeout=2)
|
|
second.join(timeout=2)
|
|
assert errors == []
|
|
db.close()
|
|
|
|
|
|
def test_postgres_schema_is_current_idempotent_and_readers_are_read_only(tmp_path: Path):
|
|
db = Database(tmp_path / "schema-postgres")
|
|
db.initialize()
|
|
with db.transaction() as conn:
|
|
conn.executemany(
|
|
"INSERT INTO jobs(id,kind,payload_json,status,priority,run_after,created_at) VALUES(?,?,?,?,?,?,?)",
|
|
(
|
|
("legacy-audio-queued", "transcribe_audio", "{}", "queued", 30, "now", "now"),
|
|
("legacy-audio-complete", "transcribe_audio", "{}", "completed", 30, "now", "now"),
|
|
("legacy-index-queued", "index_video", "{}", "queued", 20, "now", "now"),
|
|
("legacy-index-complete", "index_video", "{}", "completed", 20, "now", "now"),
|
|
),
|
|
)
|
|
db.initialize()
|
|
|
|
assert db.setting("schema_version") == SCHEMA_VERSION
|
|
with db.read() as reader:
|
|
assert reader.execute("SELECT extversion FROM pg_extension WHERE extname='vector'").fetchone()[0]
|
|
upload_columns = {
|
|
row[0]
|
|
for row in reader.execute(
|
|
"SELECT column_name FROM information_schema.columns WHERE table_name='uploads'"
|
|
).fetchall()
|
|
}
|
|
assert {"title", "target_key", "webdav_path", "content_sha256_verified"} <= upload_columns
|
|
with db.read() as reader:
|
|
priorities = {
|
|
row["id"]: row["priority"]
|
|
for row in reader.execute(
|
|
"SELECT id,priority FROM jobs WHERE id IN ("
|
|
"'legacy-audio-queued','legacy-audio-complete','legacy-index-queued','legacy-index-complete')"
|
|
).fetchall()
|
|
}
|
|
assert reader.execute(
|
|
"SELECT 1 FROM pg_indexes WHERE indexname='idx_text_entries_raw_text_trgm'"
|
|
).fetchone()
|
|
assert priorities == {
|
|
"legacy-audio-queued": 20,
|
|
"legacy-audio-complete": 30,
|
|
"legacy-index-queued": 10,
|
|
"legacy-index-complete": 20,
|
|
}
|
|
|
|
with db.read() as reader:
|
|
reader_pid = reader.execute("SELECT pg_backend_pid()").fetchone()[0]
|
|
observer = db.connect()
|
|
try:
|
|
state = observer.execute(
|
|
"SELECT state FROM pg_stat_activity WHERE pid=?", (reader_pid,)
|
|
).fetchone()[0]
|
|
finally:
|
|
observer.rollback()
|
|
observer.close()
|
|
assert state == "idle"
|
|
|
|
with pytest.raises(Exception, match="read-only"):
|
|
with db.read() as reader:
|
|
reader.execute("INSERT INTO settings(key,value,updated_at) VALUES('forbidden','1','now')")
|
|
|
|
|
|
def test_schema_initialization_backfills_job_resources_with_legacy_postgres_json_syntax(tmp_path: Path):
|
|
db = Database(tmp_path / "job-resource-backfill-postgres")
|
|
db.initialize()
|
|
with db.transaction() as conn:
|
|
conn.executemany(
|
|
"INSERT INTO jobs(id,kind,payload_json,run_after,created_at) VALUES(?,?,?,?,?)",
|
|
(
|
|
("single-job", "index_video", '{"video_id":"video-one"}', "now", "now"),
|
|
(
|
|
"bulk-job",
|
|
"bulk_index",
|
|
'{"video_ids":["video-two","video-three","video-two"]}',
|
|
"now",
|
|
"now",
|
|
),
|
|
("unrelated-job", "scan_source", '{"source_id":"source-one"}', "now", "now"),
|
|
),
|
|
)
|
|
|
|
# Re-running initialize simulates an installation/upgrade which needs to
|
|
# populate the resource relation for jobs created before that table existed.
|
|
# The implementation deliberately uses jsonb casts/operators available in
|
|
# older supported PostgreSQL releases and must not require the PG16-only
|
|
# SQL/JSON ``IS JSON`` predicate.
|
|
db.initialize()
|
|
|
|
with db.read() as conn:
|
|
resources = {
|
|
(row["job_id"], row["resource_type"], row["resource_id"])
|
|
for row in conn.execute(
|
|
"SELECT job_id,resource_type,resource_id FROM job_resources ORDER BY job_id,resource_id"
|
|
).fetchall()
|
|
}
|
|
assert resources == {
|
|
("single-job", "video", "video-one"),
|
|
("bulk-job", "video", "video-two"),
|
|
("bulk-job", "video", "video-three"),
|
|
}
|
|
|
|
|
|
def test_gateway_media_tokens_follow_session_lifecycle(tmp_path: Path):
|
|
db = Database(tmp_path / "media-token-postgres")
|
|
db.initialize()
|
|
auth = AuthService(db, session_days=1)
|
|
auth.setup("gateway media token test password")
|
|
session_token, _, _ = auth.login_gateway("nas-admin", "Administrator")
|
|
media_token, expires = auth.create_gateway_media_token(session_token)
|
|
|
|
assert expires.isoformat()
|
|
assert auth.gateway_media_session(media_token, "nas-admin")["auth_source"] == "gateway"
|
|
assert auth.gateway_media_session(media_token, "another-admin") is None
|
|
|
|
auth.logout(session_token)
|
|
assert auth.gateway_media_session(media_token, "nas-admin") is None
|
|
with db.read() as conn:
|
|
assert conn.execute("SELECT count(*) FROM gateway_media_tokens").fetchone()[0] == 0
|