Merge branch 'main' into copilot/add-document-sharing-feature

This commit is contained in:
Christian Krakau-Louis
2026-03-08 23:12:27 +01:00
committed by GitHub
31 changed files with 5345 additions and 28 deletions
+671
View File
@@ -0,0 +1,671 @@
"""Unit tests for app/cli.py — DocuElevate CLI tool.
Tests cover:
- Root command group option handling (URL, token, format)
- list command with various filters
- upload command (single file, batch, error handling)
- download command (with/without --output, Content-Disposition parsing)
- search command with filters
- token sub-commands (create, list, revoke)
- Helper functions (_build_headers, _api, _require_ok, _output, _print_table)
- Environment variable configuration
- Missing-token error handling
"""
import json
from unittest.mock import MagicMock, patch
import pytest
import requests as req_module
from click.testing import CliRunner
from app.cli import (
_api,
_build_headers,
_output,
_print_table,
_require_ok,
cli,
main,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_response(status_code: int = 200, json_data=None, text: str = "", headers: dict | None = None):
"""Create a mock requests.Response."""
mock = MagicMock(spec=req_module.Response)
mock.status_code = status_code
mock.text = text
mock.headers = headers or {}
if json_data is not None:
mock.json.return_value = json_data
else:
mock.json.side_effect = ValueError("No JSON")
return mock
# ---------------------------------------------------------------------------
# Unit tests for helper functions
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestBuildHeaders:
def test_returns_authorization_header(self):
headers = _build_headers("de_mytoken")
assert headers == {"Authorization": "Bearer de_mytoken"}
@pytest.mark.unit
class TestApi:
def test_successful_request(self):
mock_resp = _make_response(200, json_data={"ok": True})
with patch("app.cli.requests.request", return_value=mock_resp) as mock_req:
resp = _api("GET", "http://localhost:8000", "/api/files", "de_tok")
mock_req.assert_called_once()
call_kwargs = mock_req.call_args
assert call_kwargs[0][0] == "GET"
assert call_kwargs[0][1] == "http://localhost:8000/api/files"
assert resp.status_code == 200
def test_strips_trailing_slash_from_base_url(self):
mock_resp = _make_response(200, json_data={})
with patch("app.cli.requests.request", return_value=mock_resp) as mock_req:
_api("GET", "http://localhost:8000/", "/api/files", "de_tok")
assert mock_req.call_args[0][1] == "http://localhost:8000/api/files"
def test_connection_error_raises_click_exception(self):
import click
with patch("app.cli.requests.request", side_effect=req_module.ConnectionError("refused")):
with pytest.raises(click.ClickException, match="Could not connect"):
_api("GET", "http://localhost:8000", "/api/files", "de_tok")
def test_timeout_raises_click_exception(self):
import click
with patch("app.cli.requests.request", side_effect=req_module.Timeout("timed out")):
with pytest.raises(click.ClickException, match="timed out"):
_api("GET", "http://localhost:8000", "/api/files", "de_tok")
@pytest.mark.unit
class TestRequireOk:
def test_returns_json_on_success(self):
mock_resp = _make_response(200, json_data={"data": [1, 2, 3]})
result = _require_ok(mock_resp)
assert result == {"data": [1, 2, 3]}
def test_raises_on_400(self):
import click
mock_resp = _make_response(400, json_data={"detail": "Bad request"})
with pytest.raises(click.ClickException, match="API error 400"):
_require_ok(mock_resp)
def test_raises_on_404(self):
import click
mock_resp = _make_response(404, json_data={"detail": "Not found"})
with pytest.raises(click.ClickException, match="404"):
_require_ok(mock_resp)
def test_raises_on_500_with_text_fallback(self):
import click
mock_resp = _make_response(500, text="Internal Server Error")
mock_resp.json.side_effect = ValueError("no json")
with pytest.raises(click.ClickException, match="500"):
_require_ok(mock_resp)
def test_returns_empty_dict_when_no_json(self):
mock_resp = _make_response(200)
mock_resp.json.side_effect = ValueError("no json")
result = _require_ok(mock_resp)
assert result == {}
@pytest.mark.unit
class TestOutput:
def test_json_format(self, capsys):
_output({"key": "value"}, "json")
captured = capsys.readouterr()
parsed = json.loads(captured.out)
assert parsed == {"key": "value"}
def test_table_format_dict(self, capsys):
_output({"id": 1, "name": "test"}, "table")
captured = capsys.readouterr()
assert "id" in captured.out
assert "name" in captured.out
def test_table_format_list(self, capsys):
_output([{"id": 1, "name": "file1"}, {"id": 2, "name": "file2"}], "table")
captured = capsys.readouterr()
assert "file1" in captured.out
assert "file2" in captured.out
def test_table_empty_list(self, capsys):
_output([], "table")
# Should not raise, output can be empty or a JSON representation
capsys.readouterr()
def test_table_non_dict_items(self, capsys):
_output(["item1", "item2"], "table")
captured = capsys.readouterr()
assert "item1" in captured.out
@pytest.mark.unit
class TestPrintTable:
def test_single_dict(self, capsys):
_print_table({"id": 42, "name": "doc"})
captured = capsys.readouterr()
assert "42" in captured.out
assert "doc" in captured.out
def test_list_of_dicts(self, capsys):
_print_table([{"id": 1, "name": "a"}, {"id": 2, "name": "bb"}])
captured = capsys.readouterr()
assert "ID" in captured.out
assert "NAME" in captured.out
assert "a" in captured.out
assert "bb" in captured.out
def test_fallback_json_for_non_dict_list_items(self, capsys):
_print_table([1, 2, 3])
captured = capsys.readouterr()
assert "1" in captured.out
def test_fallback_json_for_scalar(self, capsys):
_print_table("plain string")
capsys.readouterr() # just assert no exception
# ---------------------------------------------------------------------------
# CLI integration tests via CliRunner
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestMissingToken:
"""Commands must fail gracefully when no token is supplied."""
def test_list_without_token(self):
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "list"])
assert result.exit_code != 0
assert "DOCUELEVATE_API_TOKEN" in result.output or "No API token" in result.output
def test_upload_without_token(self, tmp_path):
f = tmp_path / "test.pdf"
f.write_bytes(b"%PDF-1.4")
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "upload", str(f)])
assert result.exit_code != 0
def test_search_without_token(self):
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "search", "invoice"])
assert result.exit_code != 0
def test_token_create_without_token(self):
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "token", "create", "test"])
assert result.exit_code != 0
def test_token_list_without_token(self):
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "token", "list"])
assert result.exit_code != 0
def test_token_revoke_without_token(self):
runner = CliRunner()
result = runner.invoke(cli, ["--url", "http://localhost:8000", "token", "revoke", "--yes", "1"])
assert result.exit_code != 0
@pytest.mark.unit
class TestListCommand:
def test_list_success_table(self):
files_data = {
"files": [
{
"id": 1,
"original_filename": "test.pdf",
"file_size": 1024,
"status": "completed",
"created_at": "2026-01-01T00:00:00",
},
],
"pagination": {"page": 1, "pages": 1, "total": 1},
}
mock_resp = _make_response(200, json_data=files_data)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "list"])
assert result.exit_code == 0
assert "test.pdf" in result.output
def test_list_success_json(self):
files_data = {
"files": [{"id": 1, "original_filename": "file.pdf"}],
"pagination": {"page": 1, "pages": 1, "total": 1},
}
mock_resp = _make_response(200, json_data=files_data)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "list"])
assert result.exit_code == 0
parsed = json.loads(result.output)
assert isinstance(parsed, list)
assert parsed[0]["id"] == 1
def test_list_with_filters(self):
mock_resp = _make_response(200, json_data={"files": [], "pagination": {"page": 1, "pages": 0, "total": 0}})
with patch("app.cli._api", return_value=mock_resp) as mock_api:
runner = CliRunner()
result = runner.invoke(
cli,
["--token", "de_tok", "list", "--status", "completed", "--mime-type", "application/pdf"],
)
assert result.exit_code == 0
call_kwargs = mock_api.call_args[1]
assert call_kwargs["params"]["status"] == "completed"
assert call_kwargs["params"]["mime_type"] == "application/pdf"
def test_list_api_error(self):
mock_resp = _make_response(401, json_data={"detail": "Unauthorized"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_bad", "list"])
assert result.exit_code != 0
assert "401" in result.output
def test_list_raw_list_response(self):
"""Handles when the API returns a plain list (not paginated dict)."""
files_data = [{"id": 1, "original_filename": "a.pdf"}]
mock_resp = _make_response(200, json_data=files_data)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "list"])
assert result.exit_code == 0
@pytest.mark.unit
class TestUploadCommand:
def test_upload_single_file_success(self, tmp_path):
f = tmp_path / "report.pdf"
f.write_bytes(b"%PDF-1.4 content")
mock_resp = _make_response(201, json_data={"task_id": "abc-123", "status": "queued"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "upload", str(f)])
assert result.exit_code == 0
assert "abc-123" in result.output
def test_upload_multiple_files_success(self, tmp_path):
files = []
for i in range(3):
f = tmp_path / f"file{i}.pdf"
f.write_bytes(b"PDF")
files.append(str(f))
mock_resp = _make_response(201, json_data={"task_id": f"task-{0}", "status": "queued"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "upload", *files])
assert result.exit_code == 0
def test_upload_single_file_api_error(self, tmp_path):
f = tmp_path / "bad.pdf"
f.write_bytes(b"data")
mock_resp = _make_response(413, json_data={"detail": "File too large"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "upload", str(f)])
assert result.exit_code == 1
assert "failed" in result.output.lower() or "error" in result.output.lower()
def test_upload_json_output(self, tmp_path):
f = tmp_path / "test.pdf"
f.write_bytes(b"PDF")
mock_resp = _make_response(201, json_data={"task_id": "t1", "status": "queued"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "upload", str(f)])
assert result.exit_code == 0
# Progress lines (stderr) are mixed with JSON stdout in CliRunner.
# The JSON array is the last block in the output starting with '['.
import re
json_match = re.search(r"(\[\s*\{.*?\}\s*\])", result.output, re.DOTALL)
assert json_match is not None, f"No JSON array found in: {result.output!r}"
parsed = json.loads(json_match.group(1))
assert isinstance(parsed, list)
assert parsed[0]["status"] == "queued"
def test_upload_partial_failure(self, tmp_path):
"""Mixed success/failure: exit code 1 if any upload fails."""
f1 = tmp_path / "ok.pdf"
f1.write_bytes(b"PDF")
f2 = tmp_path / "fail.pdf"
f2.write_bytes(b"PDF")
ok_resp = _make_response(201, json_data={"task_id": "t1"})
err_resp = _make_response(500, json_data={"detail": "Server error"})
with patch("app.cli._api", side_effect=[ok_resp, err_resp]):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "upload", str(f1), str(f2)])
assert result.exit_code == 1
@pytest.mark.unit
class TestDownloadCommand:
def test_download_with_explicit_output(self, tmp_path):
dest = tmp_path / "out.pdf"
mock_resp = _make_response(200, headers={"content-disposition": 'attachment; filename="doc.pdf"'})
mock_resp.iter_content.return_value = [b"PDF content"]
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
with runner.isolated_filesystem():
result = runner.invoke(
cli,
["--token", "de_tok", "download", "42", "--output", str(dest)],
)
assert result.exit_code == 0
assert dest.exists()
def test_download_filename_from_content_disposition(self, tmp_path):
mock_resp = _make_response(200, headers={"content-disposition": 'attachment; filename="invoice.pdf"'})
mock_resp.iter_content.return_value = [b"PDF data"]
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
with runner.isolated_filesystem():
result = runner.invoke(cli, ["--token", "de_tok", "download", "7"])
assert result.exit_code == 0
assert "invoice.pdf" in result.output
def test_download_filename_from_content_disposition_utf8(self, tmp_path):
mock_resp = _make_response(
200,
headers={"content-disposition": "attachment; filename*=UTF-8''Rechnung%202026.pdf"},
)
mock_resp.iter_content.return_value = [b"PDF data"]
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
with runner.isolated_filesystem():
result = runner.invoke(cli, ["--token", "de_tok", "download", "8"])
assert result.exit_code == 0
assert "Rechnung" in result.output
def test_download_fallback_filename(self):
mock_resp = _make_response(200, headers={"content-disposition": ""})
mock_resp.iter_content.return_value = [b"data"]
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
with runner.isolated_filesystem():
result = runner.invoke(cli, ["--token", "de_tok", "download", "99"])
assert result.exit_code == 0
assert "file_99" in result.output
def test_download_api_error(self):
mock_resp = _make_response(404, json_data={"detail": "Not found"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "download", "999"])
assert result.exit_code != 0
assert "404" in result.output
def test_download_original_version(self, tmp_path):
mock_resp = _make_response(200, headers={"content-disposition": 'attachment; filename="orig.pdf"'})
mock_resp.iter_content.return_value = [b"original"]
with patch("app.cli._api", return_value=mock_resp) as mock_api:
runner = CliRunner()
with runner.isolated_filesystem():
result = runner.invoke(cli, ["--token", "de_tok", "download", "5", "--version", "original"])
assert result.exit_code == 0
call_kwargs = mock_api.call_args[1]
assert call_kwargs["params"]["version"] == "original"
@pytest.mark.unit
class TestSearchCommand:
def test_search_success_table(self):
payload = {
"results": [
{
"file_id": 1,
"original_filename": "inv.pdf",
"document_type": "Invoice",
"tags": ["amazon"],
}
],
"total": 1,
"pages": 1,
}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "search", "invoice"])
assert result.exit_code == 0
assert "inv.pdf" in result.output
def test_search_success_json(self):
payload = {"results": [{"file_id": 2}], "total": 1, "pages": 1}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "search", "test"])
assert result.exit_code == 0
parsed = json.loads(result.output)
assert isinstance(parsed, list)
assert parsed[0]["file_id"] == 2
def test_search_with_filters_passed_to_api(self):
payload = {"results": [], "total": 0, "pages": 0}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp) as mock_api:
runner = CliRunner()
result = runner.invoke(
cli,
[
"--token",
"de_tok",
"search",
"contract",
"--document-type",
"Contract",
"--tags",
"legal",
"--language",
"en",
"--mime-type",
"application/pdf",
],
)
assert result.exit_code == 0
params = mock_api.call_args[1]["params"]
assert params["document_type"] == "Contract"
assert params["tags"] == "legal"
assert params["language"] == "en"
assert params["mime_type"] == "application/pdf"
def test_search_api_error(self):
mock_resp = _make_response(400, json_data={"detail": "Invalid query"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "search", "bad"])
assert result.exit_code != 0
def test_search_plain_list_response(self):
"""Handles when API returns a plain list."""
mock_resp = _make_response(200, json_data=[{"file_id": 3}])
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "search", "x"])
assert result.exit_code == 0
@pytest.mark.unit
class TestTokenCreate:
def test_create_token_table(self):
payload = {
"id": 5,
"name": "CI Pipeline",
"token_prefix": "de_Abc123",
"token": "de_Abc123_fulltoken",
"is_active": True,
"last_used_at": None,
"last_used_ip": None,
"created_at": "2026-01-01T00:00:00",
"revoked_at": None,
}
mock_resp = _make_response(201, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "create", "CI Pipeline"])
assert result.exit_code == 0
assert "de_Abc123_fulltoken" in result.output
assert "CI Pipeline" in result.output
def test_create_token_json(self):
payload = {"id": 6, "name": "Script", "token": "de_full", "token_prefix": "de_fu"}
mock_resp = _make_response(201, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "token", "create", "Script"])
assert result.exit_code == 0
parsed = json.loads(result.output)
assert parsed["token"] == "de_full"
def test_create_token_api_error(self):
mock_resp = _make_response(422, json_data={"detail": "name too short"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "create", "x"])
assert result.exit_code != 0
def test_create_token_unexpected_response_format(self):
"""If API returns a list instead of dict, should fail gracefully."""
mock_resp = _make_response(201, json_data=[{"id": 1}])
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "create", "bad"])
assert result.exit_code != 0
@pytest.mark.unit
class TestTokenList:
def test_list_tokens_table(self):
payload = [
{
"id": 1,
"name": "CI",
"token_prefix": "de_Ab",
"is_active": True,
"last_used_at": "2026-01-15T10:00:00",
"last_used_ip": "10.0.0.1",
"created_at": "2026-01-01T00:00:00",
"revoked_at": None,
}
]
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "list"])
assert result.exit_code == 0
assert "CI" in result.output
def test_list_tokens_json(self):
payload = [{"id": 2, "name": "S", "is_active": False}]
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "--format", "json", "token", "list"])
assert result.exit_code == 0
parsed = json.loads(result.output)
assert parsed[0]["id"] == 2
def test_list_tokens_unexpected_format(self):
"""If API returns a dict instead of list, should fail gracefully."""
mock_resp = _make_response(200, json_data={"id": 1})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "list"])
assert result.exit_code != 0
@pytest.mark.unit
class TestTokenRevoke:
def test_revoke_with_yes_flag(self):
mock_resp = _make_response(200, json_data={"detail": "Token revoked"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "revoke", "--yes", "3"])
assert result.exit_code == 0
assert "revoked" in result.output.lower()
def test_revoke_prompts_for_confirmation(self):
mock_resp = _make_response(200, json_data={"detail": "Token revoked"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "revoke", "3"], input="y\n")
assert result.exit_code == 0
def test_revoke_aborts_on_no(self):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "revoke", "3"], input="n\n")
assert result.exit_code != 0
def test_revoke_api_error(self):
mock_resp = _make_response(404, json_data={"detail": "Token not found"})
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner()
result = runner.invoke(cli, ["--token", "de_tok", "token", "revoke", "--yes", "999"])
assert result.exit_code != 0
assert "404" in result.output
@pytest.mark.unit
class TestEnvironmentVariables:
def test_token_from_env_var(self):
payload = {"files": [], "pagination": {"page": 1, "pages": 0, "total": 0}}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp):
runner = CliRunner(env={"DOCUELEVATE_API_TOKEN": "de_envtoken"})
result = runner.invoke(cli, ["list"])
assert result.exit_code == 0
def test_url_from_env_var(self):
payload = {"files": [], "pagination": {"page": 1, "pages": 0, "total": 0}}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp) as mock_api:
runner = CliRunner(env={"DOCUELEVATE_URL": "http://my-server:9000", "DOCUELEVATE_API_TOKEN": "de_tok"})
result = runner.invoke(cli, ["list"])
assert result.exit_code == 0
assert mock_api.call_args[0][1] == "http://my-server:9000"
def test_explicit_token_overrides_env(self):
payload = {"files": [], "pagination": {"page": 1, "pages": 0, "total": 0}}
mock_resp = _make_response(200, json_data=payload)
with patch("app.cli._api", return_value=mock_resp) as mock_api:
runner = CliRunner(env={"DOCUELEVATE_API_TOKEN": "de_env"})
result = runner.invoke(cli, ["--token", "de_explicit", "list"])
assert result.exit_code == 0
# Token passed to _api should be the explicit one
token_arg = mock_api.call_args[0][3]
assert token_arg == "de_explicit"
@pytest.mark.unit
class TestMainEntryPoint:
def test_main_invokes_cli(self):
"""main() should be callable without errors (help flag)."""
runner = CliRunner()
with patch("app.cli.cli") as mock_cli:
main()
mock_cli.assert_called_once()
+832
View File
@@ -0,0 +1,832 @@
"""Tests for the per-user notification system (app/api/notifications.py)."""
import json
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from app.database import Base, get_db
from app.models import InAppNotification, UserNotificationPreference, UserNotificationTarget
# ---------------------------------------------------------------------------
# Test data
# ---------------------------------------------------------------------------
_OWNER = "notifuser@example.com"
_OTHER_OWNER = "other@example.com"
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture()
def notif_engine():
"""In-memory SQLite engine for notification tests."""
engine = create_engine(
"sqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
yield engine
Base.metadata.drop_all(bind=engine)
@pytest.fixture()
def notif_session(notif_engine):
"""DB session scoped to one test."""
Session = sessionmaker(bind=notif_engine)
session = Session()
yield session
session.close()
def _make_client(notif_engine, owner_id: str = _OWNER) -> TestClient:
"""Return a TestClient with *owner_id* injected as the authenticated user."""
from app.api.notifications import _get_owner_id
from app.main import app
Session = sessionmaker(bind=notif_engine)
def _override_get_db():
session = Session()
try:
yield session
finally:
session.close()
def _override_owner():
return owner_id
app.dependency_overrides[get_db] = _override_get_db
app.dependency_overrides[_get_owner_id] = _override_owner
return TestClient(app, base_url="http://localhost", raise_server_exceptions=False)
def _cleanup(app):
"""Remove dependency overrides after test."""
app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# Tests – Auth / 401 guard
# ---------------------------------------------------------------------------
class TestAuthGuard:
"""Verify that unauthenticated requests are rejected."""
@pytest.mark.unit
def test_inbox_requires_auth(self):
"""GET /api/user-notifications/inbox should return 401 when not authenticated."""
from app.main import app
client = TestClient(app, base_url="http://localhost", raise_server_exceptions=False)
resp = client.get("/api/user-notifications/inbox")
assert resp.status_code == 401
@pytest.mark.unit
def test_unread_count_requires_auth(self):
"""GET /api/user-notifications/inbox/unread-count should return 401 when not authenticated."""
from app.main import app
client = TestClient(app, base_url="http://localhost", raise_server_exceptions=False)
resp = client.get("/api/user-notifications/inbox/unread-count")
assert resp.status_code == 401
# ---------------------------------------------------------------------------
# Tests – Inbox
# ---------------------------------------------------------------------------
class TestInbox:
"""Tests for the in-app notification inbox."""
@pytest.mark.unit
def test_unread_count_empty(self, notif_engine):
"""Unread count should be 0 when no notifications exist."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/inbox/unread-count")
assert resp.status_code == 200
assert resp.json() == {"count": 0}
finally:
_cleanup(app)
@pytest.mark.unit
def test_inbox_empty(self, notif_engine):
"""Listing inbox when empty should return an empty list."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/inbox")
assert resp.status_code == 200
assert resp.json() == []
finally:
_cleanup(app)
@pytest.mark.unit
def test_inbox_shows_notifications(self, notif_engine, notif_session):
"""Inbox should return notifications for the authenticated user."""
from app.main import app
notif_session.add(
InAppNotification(
owner_id=_OWNER,
event_type="document.processed",
title="Test",
message="Done",
)
)
notif_session.commit()
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/inbox")
assert resp.status_code == 200
items = resp.json()
assert len(items) == 1
assert items[0]["title"] == "Test"
assert items[0]["is_read"] is False
finally:
_cleanup(app)
@pytest.mark.unit
def test_inbox_isolation(self, notif_engine, notif_session):
"""Users should only see their own notifications."""
from app.main import app
notif_session.add(
InAppNotification(
owner_id=_OTHER_OWNER,
event_type="document.processed",
title="Other user notif",
message="Not yours",
)
)
notif_session.commit()
client = _make_client(notif_engine, _OWNER)
try:
resp = client.get("/api/user-notifications/inbox")
assert resp.status_code == 200
assert resp.json() == []
finally:
_cleanup(app)
@pytest.mark.unit
def test_unread_count_reflects_notifications(self, notif_engine, notif_session):
"""Unread count should reflect actual unread notifications."""
from app.main import app
for i in range(3):
notif_session.add(
InAppNotification(
owner_id=_OWNER,
event_type="document.processed",
title=f"Notif {i}",
message="",
is_read=False,
)
)
notif_session.commit()
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/inbox/unread-count")
assert resp.status_code == 200
assert resp.json()["count"] == 3
finally:
_cleanup(app)
@pytest.mark.unit
def test_mark_read(self, notif_engine, notif_session):
"""Marking a notification as read should update is_read."""
from app.main import app
notif = InAppNotification(
owner_id=_OWNER,
event_type="document.processed",
title="Unread",
message="",
)
notif_session.add(notif)
notif_session.commit()
notif_session.refresh(notif)
notif_id = notif.id
client = _make_client(notif_engine)
try:
resp = client.post(f"/api/user-notifications/inbox/{notif_id}/read")
assert resp.status_code == 200
# Verify in DB
notif_session.refresh(notif)
assert notif.is_read is True
finally:
_cleanup(app)
@pytest.mark.unit
def test_mark_read_wrong_user(self, notif_engine, notif_session):
"""Marking another user's notification should return 404."""
from app.main import app
notif = InAppNotification(
owner_id=_OTHER_OWNER,
event_type="document.processed",
title="Other",
message="",
)
notif_session.add(notif)
notif_session.commit()
notif_session.refresh(notif)
notif_id = notif.id
client = _make_client(notif_engine, _OWNER)
try:
resp = client.post(f"/api/user-notifications/inbox/{notif_id}/read")
assert resp.status_code == 404
finally:
_cleanup(app)
@pytest.mark.unit
def test_mark_all_read(self, notif_engine, notif_session):
"""Mark all read should set all user's notifications to read."""
from app.main import app
for i in range(4):
notif_session.add(
InAppNotification(
owner_id=_OWNER,
event_type="document.processed",
title=f"N{i}",
message="",
is_read=False,
)
)
notif_session.commit()
client = _make_client(notif_engine)
try:
resp = client.post("/api/user-notifications/inbox/read-all")
assert resp.status_code == 200
count_resp = client.get("/api/user-notifications/inbox/unread-count")
assert count_resp.json()["count"] == 0
finally:
_cleanup(app)
# ---------------------------------------------------------------------------
# Tests – Notification Targets
# ---------------------------------------------------------------------------
class TestTargets:
"""Tests for notification target CRUD."""
@pytest.mark.unit
def test_list_targets_empty(self, notif_engine):
"""Listing targets when none exist should return empty list."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/targets")
assert resp.status_code == 200
assert resp.json() == []
finally:
_cleanup(app)
@pytest.mark.unit
def test_create_email_target(self, notif_engine):
"""Creating an email target should persist and mask the password in response."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.post(
"/api/user-notifications/targets",
json={
"channel_type": "email",
"name": "My Gmail",
"config": {
"smtp_host": "smtp.gmail.com",
"smtp_port": 587,
"smtp_username": "me@gmail.com",
"smtp_password": "s3cr3t",
"recipient_email": "me@gmail.com",
"smtp_use_tls": True,
},
"is_active": True,
},
)
assert resp.status_code == 201, resp.text
data = resp.json()
assert data["channel_type"] == "email"
assert data["name"] == "My Gmail"
assert data["is_active"] is True
# Password must be masked
assert data["config"]["smtp_password"] == "****"
finally:
_cleanup(app)
@pytest.mark.unit
def test_create_webhook_target(self, notif_engine):
"""Creating a webhook target should persist correctly."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.post(
"/api/user-notifications/targets",
json={
"channel_type": "webhook",
"name": "Slack Webhook",
"config": {"url": "https://hooks.slack.com/abc", "secret": ""},
"is_active": True,
},
)
assert resp.status_code == 201, resp.text
data = resp.json()
assert data["channel_type"] == "webhook"
assert data["name"] == "Slack Webhook"
finally:
_cleanup(app)
@pytest.mark.unit
def test_create_target_invalid_channel_type(self, notif_engine):
"""Creating a target with an invalid channel_type should return 422."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.post(
"/api/user-notifications/targets",
json={"channel_type": "sms", "name": "Bad", "config": {}},
)
assert resp.status_code == 422
finally:
_cleanup(app)
@pytest.mark.unit
def test_list_targets_returns_created(self, notif_engine):
"""Listing targets should include newly created ones."""
from app.main import app
client = _make_client(notif_engine)
try:
client.post(
"/api/user-notifications/targets",
json={"channel_type": "webhook", "name": "W1", "config": {"url": "https://example.com"}},
)
client.post(
"/api/user-notifications/targets",
json={"channel_type": "email", "name": "E1", "config": {"smtp_host": "smtp.example.com"}},
)
resp = client.get("/api/user-notifications/targets")
assert resp.status_code == 200
assert len(resp.json()) == 2
finally:
_cleanup(app)
@pytest.mark.unit
def test_targets_isolation(self, notif_engine):
"""Users should only see their own targets."""
from app.main import app
client_a = _make_client(notif_engine, _OWNER)
try:
client_a.post(
"/api/user-notifications/targets",
json={"channel_type": "webhook", "name": "Owner A target", "config": {"url": "https://a.example.com"}},
)
finally:
_cleanup(app)
client_b = _make_client(notif_engine, _OTHER_OWNER)
try:
resp = client_b.get("/api/user-notifications/targets")
assert resp.status_code == 200
assert resp.json() == []
finally:
_cleanup(app)
@pytest.mark.unit
def test_update_target(self, notif_engine):
"""Updating a target should change its name and active status."""
from app.main import app
client = _make_client(notif_engine)
try:
create_resp = client.post(
"/api/user-notifications/targets",
json={"channel_type": "webhook", "name": "Old Name", "config": {"url": "https://x.com"}},
)
target_id = create_resp.json()["id"]
resp = client.put(
f"/api/user-notifications/targets/{target_id}",
json={"name": "New Name", "is_active": False},
)
assert resp.status_code == 200
data = resp.json()
assert data["name"] == "New Name"
assert data["is_active"] is False
finally:
_cleanup(app)
@pytest.mark.unit
def test_update_target_wrong_user(self, notif_engine, notif_session):
"""Updating another user's target should return 404."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OTHER_OWNER,
channel_type="webhook",
name="Other target",
config=json.dumps({"url": "https://other.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
client = _make_client(notif_engine, _OWNER)
try:
resp = client.put(
f"/api/user-notifications/targets/{target.id}",
json={"name": "Hacked"},
)
assert resp.status_code == 404
finally:
_cleanup(app)
@pytest.mark.unit
def test_delete_target(self, notif_engine):
"""Deleting a target should remove it from the list."""
from app.main import app
client = _make_client(notif_engine)
try:
create_resp = client.post(
"/api/user-notifications/targets",
json={"channel_type": "webhook", "name": "To Delete", "config": {"url": "https://x.com"}},
)
target_id = create_resp.json()["id"]
del_resp = client.delete(f"/api/user-notifications/targets/{target_id}")
assert del_resp.status_code == 200
list_resp = client.get("/api/user-notifications/targets")
assert list_resp.json() == []
finally:
_cleanup(app)
@pytest.mark.unit
def test_delete_target_wrong_user(self, notif_engine, notif_session):
"""Deleting another user's target should return 404."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OTHER_OWNER,
channel_type="webhook",
name="Not yours",
config=json.dumps({"url": "https://other.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
client = _make_client(notif_engine, _OWNER)
try:
resp = client.delete(f"/api/user-notifications/targets/{target.id}")
assert resp.status_code == 404
finally:
_cleanup(app)
@pytest.mark.unit
def test_delete_target_also_removes_preferences(self, notif_engine, notif_session):
"""Deleting a target should also remove associated preferences."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OWNER,
channel_type="webhook",
name="With prefs",
config=json.dumps({"url": "https://x.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
pref = UserNotificationPreference(
owner_id=_OWNER,
event_type="document.processed",
channel_type="webhook",
target_id=target.id,
is_enabled=True,
)
notif_session.add(pref)
notif_session.commit()
client = _make_client(notif_engine, _OWNER)
try:
resp = client.delete(f"/api/user-notifications/targets/{target.id}")
assert resp.status_code == 200
remaining = (
notif_session.query(UserNotificationPreference)
.filter(UserNotificationPreference.owner_id == _OWNER)
.all()
)
assert remaining == []
finally:
_cleanup(app)
# ---------------------------------------------------------------------------
# Tests – Preferences
# ---------------------------------------------------------------------------
class TestPreferences:
"""Tests for notification preferences CRUD."""
@pytest.mark.unit
def test_get_preferences_empty(self, notif_engine):
"""Getting preferences returns event_types and event_labels even with no prefs set."""
from app.main import app
client = _make_client(notif_engine)
try:
resp = client.get("/api/user-notifications/preferences")
assert resp.status_code == 200
data = resp.json()
assert "event_types" in data
assert "event_labels" in data
assert "preferences" in data
assert "document.processed" in data["event_types"]
assert "document.failed" in data["event_types"]
finally:
_cleanup(app)
@pytest.mark.unit
def test_update_preferences(self, notif_engine, notif_session):
"""Updating preferences should persist the changes."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OWNER,
channel_type="webhook",
name="My Webhook",
config=json.dumps({"url": "https://x.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
client = _make_client(notif_engine, _OWNER)
try:
resp = client.put(
"/api/user-notifications/preferences",
json={
"preferences": [
{
"event_type": "document.processed",
"channel_type": "webhook",
"is_enabled": True,
"target_id": target.id,
}
]
},
)
assert resp.status_code == 200
# Verify stored
pref = (
notif_session.query(UserNotificationPreference)
.filter(
UserNotificationPreference.owner_id == _OWNER,
UserNotificationPreference.event_type == "document.processed",
UserNotificationPreference.channel_type == "webhook",
)
.first()
)
assert pref is not None
assert pref.is_enabled is True
assert pref.target_id == target.id
finally:
_cleanup(app)
@pytest.mark.unit
def test_update_preferences_upsert(self, notif_engine, notif_session):
"""Updating preferences twice should upsert (not duplicate)."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OWNER,
channel_type="webhook",
name="W",
config=json.dumps({"url": "https://x.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
client = _make_client(notif_engine, _OWNER)
try:
pref_item = {
"event_type": "document.processed",
"channel_type": "webhook",
"is_enabled": True,
"target_id": target.id,
}
client.put("/api/user-notifications/preferences", json={"preferences": [pref_item]})
# Disable it
pref_item["is_enabled"] = False
resp = client.put("/api/user-notifications/preferences", json={"preferences": [pref_item]})
assert resp.status_code == 200
prefs = (
notif_session.query(UserNotificationPreference)
.filter(UserNotificationPreference.owner_id == _OWNER)
.all()
)
assert len(prefs) == 1
assert prefs[0].is_enabled is False
finally:
_cleanup(app)
@pytest.mark.unit
def test_update_preferences_rejects_foreign_target(self, notif_engine, notif_session):
"""Preferences referencing another user's target_id should be rejected."""
from app.main import app
other_target = UserNotificationTarget(
owner_id=_OTHER_OWNER,
channel_type="webhook",
name="Other webhook",
config=json.dumps({"url": "https://other.com"}),
)
notif_session.add(other_target)
notif_session.commit()
notif_session.refresh(other_target)
client = _make_client(notif_engine, _OWNER)
try:
resp = client.put(
"/api/user-notifications/preferences",
json={
"preferences": [
{
"event_type": "document.processed",
"channel_type": "webhook",
"is_enabled": True,
"target_id": other_target.id,
}
]
},
)
assert resp.status_code == 400
finally:
_cleanup(app)
@pytest.mark.unit
def test_get_preferences_reflects_saved(self, notif_engine, notif_session):
"""GET preferences should reflect previously saved preferences."""
from app.main import app
target = UserNotificationTarget(
owner_id=_OWNER,
channel_type="email",
name="Email target",
config=json.dumps({"smtp_host": "smtp.example.com", "recipient_email": "me@example.com"}),
)
notif_session.add(target)
notif_session.commit()
notif_session.refresh(target)
notif_session.add(
UserNotificationPreference(
owner_id=_OWNER,
event_type="document.failed",
channel_type="email",
target_id=target.id,
is_enabled=True,
)
)
notif_session.commit()
client = _make_client(notif_engine, _OWNER)
try:
resp = client.get("/api/user-notifications/preferences")
assert resp.status_code == 200
data = resp.json()
assert "document.failed" in data["preferences"]
assert "email" in data["preferences"]["document.failed"]
assert data["preferences"]["document.failed"]["email"]["is_enabled"] is True
finally:
_cleanup(app)
# ---------------------------------------------------------------------------
# Tests – user_notification service
# ---------------------------------------------------------------------------
class TestUserNotificationService:
"""Unit tests for the user notification dispatch service."""
@pytest.mark.unit
def test_create_in_app_notification(self, notif_engine, notif_session):
"""create_in_app_notification should persist a record."""
from unittest.mock import patch
from app.utils.user_notification import create_in_app_notification
Session = sessionmaker(bind=notif_engine)
with patch("app.utils.user_notification.SessionLocal", Session):
result = create_in_app_notification(
owner_id=_OWNER,
event_type="document.processed",
title="Test",
message="Done",
file_id=42,
)
assert result is not None
assert result.owner_id == _OWNER
assert result.title == "Test"
assert result.file_id == 42
@pytest.mark.unit
def test_notify_user_document_processed(self, notif_engine):
"""notify_user_document_processed should create an in-app notification."""
from unittest.mock import patch
from app.utils.user_notification import notify_user_document_processed
Session = sessionmaker(bind=notif_engine)
with patch("app.utils.user_notification.SessionLocal", Session):
notify_user_document_processed(owner_id=_OWNER, filename="test.pdf", file_id=1)
s = Session()
notifs = s.query(InAppNotification).filter(InAppNotification.owner_id == _OWNER).all()
s.close()
assert len(notifs) == 1
assert "test.pdf" in notifs[0].title
@pytest.mark.unit
def test_notify_user_document_failed(self, notif_engine):
"""notify_user_document_failed should create an in-app notification."""
from unittest.mock import patch
from app.utils.user_notification import notify_user_document_failed
Session = sessionmaker(bind=notif_engine)
with patch("app.utils.user_notification.SessionLocal", Session):
notify_user_document_failed(owner_id=_OWNER, filename="doc.pdf", error="OCR timeout")
s = Session()
notifs = s.query(InAppNotification).filter(InAppNotification.owner_id == _OWNER).all()
s.close()
assert len(notifs) == 1
assert notifs[0].event_type == "document.failed"
assert "OCR timeout" in notifs[0].message
@pytest.mark.unit
def test_send_webhook_notification_missing_url(self):
"""_send_webhook_notification should return False when url is missing."""
from app.utils.user_notification import _send_webhook_notification
result = _send_webhook_notification({}, "document.processed", "Title", "Body")
assert result is False
@pytest.mark.unit
def test_send_email_notification_missing_host(self):
"""_send_email_notification should return False when smtp_host is missing."""
from app.utils.user_notification import _send_email_notification
result = _send_email_notification({"recipient_email": "me@example.com"}, "Title", "Body")
assert result is False
@pytest.mark.unit
def test_send_email_notification_missing_recipient(self):
"""_send_email_notification should return False when recipient_email is missing."""
from app.utils.user_notification import _send_email_notification
result = _send_email_notification({"smtp_host": "smtp.example.com"}, "Title", "Body")
assert result is False
+238
View File
@@ -1057,3 +1057,241 @@ class TestMergeOCRResults:
ms.openai_model = "gpt-4"
text, _, _ = merge_ocr_results([r1, r2], "doc.pdf")
assert text == "this is the longer text from tesseract engine"
# ---------------------------------------------------------------------------
# Multi-language OCR support
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestOCRLanguageConstants:
"""Tests for the OCR_LANGUAGES constant and TESSERACT_TO_EASYOCR mapping."""
def test_ocr_languages_has_20_plus_entries(self):
"""OCR_LANGUAGES contains at least 20 language options (excluding 'auto')."""
from app.utils.ocr_provider import OCR_LANGUAGES
language_entries = {k: v for k, v in OCR_LANGUAGES.items() if v != "auto"}
assert len(language_entries) >= 20, f"Expected ≥20 languages, got {len(language_entries)}"
def test_ocr_languages_includes_auto(self):
"""OCR_LANGUAGES includes 'auto' as the first option."""
from app.utils.ocr_provider import OCR_LANGUAGES
assert "auto" in OCR_LANGUAGES.values()
def test_ocr_languages_common_languages(self):
"""OCR_LANGUAGES includes the most common European and Asian languages."""
from app.utils.ocr_provider import OCR_LANGUAGES
expected_codes = {"eng", "deu", "fra", "spa", "ita", "por", "rus", "chi_sim", "jpn", "kor"}
all_codes = set(OCR_LANGUAGES.values())
missing = expected_codes - all_codes
assert not missing, f"Missing expected language codes: {missing}"
def test_tesseract_to_easyocr_mapping(self):
"""TESSERACT_TO_EASYOCR maps common Tesseract codes to EasyOCR codes."""
from app.utils.ocr_provider import TESSERACT_TO_EASYOCR
assert TESSERACT_TO_EASYOCR["eng"] == "en"
assert TESSERACT_TO_EASYOCR["deu"] == "de"
assert TESSERACT_TO_EASYOCR["fra"] == "fr"
assert TESSERACT_TO_EASYOCR["chi_sim"] == "ch_sim"
def test_tesseract_codes_to_easyocr_single(self):
"""_tesseract_codes_to_easyocr converts a single Tesseract code."""
from app.utils.ocr_provider import _tesseract_codes_to_easyocr
result = _tesseract_codes_to_easyocr("eng")
assert result == ["en"]
def test_tesseract_codes_to_easyocr_multi(self):
"""_tesseract_codes_to_easyocr splits '+'-separated Tesseract codes."""
from app.utils.ocr_provider import _tesseract_codes_to_easyocr
result = _tesseract_codes_to_easyocr("eng+deu")
assert result == ["en", "de"]
def test_tesseract_codes_to_easyocr_passthrough_unknown(self):
"""_tesseract_codes_to_easyocr passes through codes not in the mapping."""
from app.utils.ocr_provider import _tesseract_codes_to_easyocr
# EasyOCR-native codes are passed through unchanged
result = _tesseract_codes_to_easyocr("en")
assert result == ["en"]
@pytest.mark.unit
class TestTesseractLanguageOverride:
"""Tests for per-call language override in TesseractOCRProvider."""
def test_language_override_used_in_process(self, tmp_path):
"""Language override is used instead of global setting."""
pdf = _make_pdf(tmp_path)
provider = TesseractOCRProvider(language="deu")
mock_pytesseract = Mock()
mock_pytesseract.image_to_string.return_value = "Deutsches Text"
mock_pytesseract.pytesseract = Mock()
mock_pdf2image = Mock()
mock_pdf2image.convert_from_path.return_value = [Mock()]
with (
patch.dict(
sys.modules,
{"pytesseract": mock_pytesseract, "pdf2image": mock_pdf2image},
),
patch("app.utils.ocr_provider.settings") as ms,
patch("app.utils.ocr_language_manager.ensure_tesseract_languages", return_value=[]),
):
ms.tesseract_cmd = None
ms.tesseract_language = "eng" # global setting; should be overridden
result = provider.process(pdf)
# Ensure image_to_string was called with the override language ("deu"), not global "eng"
mock_pytesseract.image_to_string.assert_called_once()
call_kwargs = mock_pytesseract.image_to_string.call_args
assert call_kwargs[1].get("lang") == "deu" or (call_kwargs[0] and call_kwargs[0][1] == "deu")
assert result.provider == "tesseract"
def test_auto_language_falls_back_to_global(self, tmp_path):
"""'auto' language override falls back to global tesseract_language setting."""
pdf = _make_pdf(tmp_path)
provider = TesseractOCRProvider(language="auto")
mock_pytesseract = Mock()
mock_pytesseract.image_to_string.return_value = ""
mock_pytesseract.pytesseract = Mock()
mock_pdf2image = Mock()
mock_pdf2image.convert_from_path.return_value = [Mock()]
with (
patch.dict(
sys.modules,
{"pytesseract": mock_pytesseract, "pdf2image": mock_pdf2image},
),
patch("app.utils.ocr_provider.settings") as ms,
patch("app.utils.ocr_language_manager.ensure_tesseract_languages", return_value=[]),
):
ms.tesseract_cmd = None
ms.tesseract_language = "fra"
provider.process(pdf)
# Should use global setting "fra" since "auto" means no override
mock_pytesseract.image_to_string.assert_called_once()
call_kwargs = mock_pytesseract.image_to_string.call_args
lang_used = call_kwargs[1].get("lang") if call_kwargs[1] else call_kwargs[0][1]
assert lang_used == "fra"
def test_none_language_falls_back_to_global(self, tmp_path):
"""None language override falls back to global setting."""
pdf = _make_pdf(tmp_path)
provider = TesseractOCRProvider(language=None)
assert provider._language_override is None
@pytest.mark.unit
class TestEasyOCRLanguageOverride:
"""Tests for per-call language override in EasyOCRProvider."""
def test_language_override_converted_and_used(self, tmp_path):
"""Tesseract-style language override is converted to EasyOCR codes."""
pdf = _make_pdf(tmp_path)
provider = EasyOCRProvider(language="deu")
mock_reader = Mock()
mock_reader.readtext.return_value = ["Deutsches Text"]
mock_easyocr = Mock()
mock_easyocr.Reader.return_value = mock_reader
mock_pdf2image = Mock()
mock_pdf2image.convert_from_path.return_value = [Mock()]
with (
patch.dict(
sys.modules,
{"easyocr": mock_easyocr, "pdf2image": mock_pdf2image},
),
patch("app.utils.ocr_provider.settings") as ms,
):
ms.easyocr_languages = "en" # global; should be overridden
ms.easyocr_gpu = False
provider.process(pdf)
# Should call Reader with ["de"] (converted from "deu"), not global ["en"]
mock_easyocr.Reader.assert_called_once()
langs_arg = mock_easyocr.Reader.call_args[0][0]
assert langs_arg == ["de"]
def test_auto_language_uses_global_setting(self, tmp_path):
"""'auto' language override falls back to global easyocr_languages setting."""
pdf = _make_pdf(tmp_path)
provider = EasyOCRProvider(language="auto")
mock_reader = Mock()
mock_reader.readtext.return_value = []
mock_easyocr = Mock()
mock_easyocr.Reader.return_value = mock_reader
mock_pdf2image = Mock()
mock_pdf2image.convert_from_path.return_value = [Mock()]
with (
patch.dict(
sys.modules,
{"easyocr": mock_easyocr, "pdf2image": mock_pdf2image},
),
patch("app.utils.ocr_provider.settings") as ms,
):
ms.easyocr_languages = "fr,es"
ms.easyocr_gpu = False
provider.process(pdf)
langs_arg = mock_easyocr.Reader.call_args[0][0]
assert langs_arg == ["fr", "es"]
@pytest.mark.unit
class TestGetOCRProvidersWithLanguage:
"""Tests for get_ocr_providers(language=...) factory."""
def test_language_passed_to_tesseract_provider(self):
"""Language override is passed to TesseractOCRProvider."""
with patch("app.utils.ocr_provider.settings") as ms:
ms.ocr_providers = "tesseract"
providers = get_ocr_providers(language="deu")
assert len(providers) == 1
assert isinstance(providers[0], TesseractOCRProvider)
assert providers[0]._language_override == "deu"
def test_language_passed_to_easyocr_provider(self):
"""Language override is passed to EasyOCRProvider."""
with patch("app.utils.ocr_provider.settings") as ms:
ms.ocr_providers = "easyocr"
providers = get_ocr_providers(language="fra")
assert len(providers) == 1
assert isinstance(providers[0], EasyOCRProvider)
assert providers[0]._language_override == "fra"
def test_language_not_passed_to_azure(self):
"""Language override is NOT passed to AzureOCRProvider (it auto-detects)."""
with patch("app.utils.ocr_provider.settings") as ms:
ms.ocr_providers = "azure"
providers = get_ocr_providers(language="deu")
assert len(providers) == 1
assert isinstance(providers[0], AzureOCRProvider)
# AzureOCRProvider has no _language_override attribute
assert not hasattr(providers[0], "_language_override")
def test_auto_language_not_passed_as_override(self):
"""'auto' language is treated as no override for Tesseract."""
with patch("app.utils.ocr_provider.settings") as ms:
ms.ocr_providers = "tesseract"
providers = get_ocr_providers(language="auto")
assert providers[0]._language_override is None
def test_none_language_no_override(self):
"""None language results in no override."""
with patch("app.utils.ocr_provider.settings") as ms:
ms.ocr_providers = "tesseract"
providers = get_ocr_providers(language=None)
assert providers[0]._language_override is None
+179
View File
@@ -842,3 +842,182 @@ startxref
mock_init_steps.assert_called_once()
called_file_id = mock_init_steps.call_args[0][1]
assert called_file_id == result["file_id"]
# ---------------------------------------------------------------------------
# _get_pipeline_ocr_language helper
# ---------------------------------------------------------------------------
@pytest.mark.unit
@pytest.mark.requires_db
def test_get_pipeline_ocr_language_returns_none_when_no_pipeline(db_session):
"""Returns None when no pipeline exists in the database."""
from app.tasks.process_document import _get_pipeline_ocr_language
# FileRecord with no pipeline_id
file_record = FileRecord(
filehash="abc123",
original_filename="test.pdf",
local_filename="/tmp/test.pdf",
file_size=1024,
mime_type="application/pdf",
is_duplicate=False,
)
db_session.add(file_record)
db_session.commit()
result = _get_pipeline_ocr_language(db_session, file_record, owner_id=None)
assert result is None
@pytest.mark.unit
@pytest.mark.requires_db
def test_get_pipeline_ocr_language_returns_language_from_system_default(db_session):
"""Returns ocr_language from the system default pipeline's OCR step config."""
import json
from app.models import Pipeline, PipelineStep
from app.tasks.process_document import _get_pipeline_ocr_language
# Create system default pipeline with OCR step configured to "deu"
pipeline = Pipeline(
owner_id=None,
name="System Default",
is_default=True,
is_active=True,
)
db_session.add(pipeline)
db_session.commit()
ocr_step = PipelineStep(
pipeline_id=pipeline.id,
position=0,
step_type="ocr",
config=json.dumps({"force_cloud_ocr": False, "ocr_language": "deu"}),
enabled=True,
)
db_session.add(ocr_step)
db_session.commit()
file_record = FileRecord(
filehash="def456",
original_filename="doc.pdf",
local_filename="/tmp/doc.pdf",
file_size=512,
mime_type="application/pdf",
is_duplicate=False,
)
db_session.add(file_record)
db_session.commit()
result = _get_pipeline_ocr_language(db_session, file_record, owner_id=None)
assert result == "deu"
@pytest.mark.unit
@pytest.mark.requires_db
def test_get_pipeline_ocr_language_auto_returns_none(db_session):
"""Returns None when ocr_language is 'auto' (should use global settings)."""
import json
from app.models import Pipeline, PipelineStep
from app.tasks.process_document import _get_pipeline_ocr_language
pipeline = Pipeline(
owner_id=None,
name="Auto Lang Pipeline",
is_default=True,
is_active=True,
)
db_session.add(pipeline)
db_session.commit()
ocr_step = PipelineStep(
pipeline_id=pipeline.id,
position=0,
step_type="ocr",
config=json.dumps({"ocr_language": "auto"}),
enabled=True,
)
db_session.add(ocr_step)
db_session.commit()
file_record = FileRecord(
filehash="ghi789",
original_filename="auto.pdf",
local_filename="/tmp/auto.pdf",
file_size=128,
mime_type="application/pdf",
is_duplicate=False,
)
db_session.add(file_record)
db_session.commit()
result = _get_pipeline_ocr_language(db_session, file_record, owner_id=None)
assert result is None
@pytest.mark.unit
@pytest.mark.requires_db
def test_get_pipeline_ocr_language_explicit_pipeline_takes_priority(db_session):
"""Explicit pipeline_id on file takes priority over system default pipeline."""
import json
from app.models import Pipeline, PipelineStep
from app.tasks.process_document import _get_pipeline_ocr_language
# System default pipeline with "eng"
sys_pipeline = Pipeline(
owner_id=None,
name="System Default",
is_default=True,
is_active=True,
)
db_session.add(sys_pipeline)
db_session.commit()
sys_step = PipelineStep(
pipeline_id=sys_pipeline.id,
position=0,
step_type="ocr",
config=json.dumps({"ocr_language": "eng"}),
enabled=True,
)
db_session.add(sys_step)
db_session.commit()
# Explicit pipeline with "fra"
explicit_pipeline = Pipeline(
owner_id="user1",
name="French Pipeline",
is_default=False,
is_active=True,
)
db_session.add(explicit_pipeline)
db_session.commit()
explicit_step = PipelineStep(
pipeline_id=explicit_pipeline.id,
position=0,
step_type="ocr",
config=json.dumps({"ocr_language": "fra"}),
enabled=True,
)
db_session.add(explicit_step)
db_session.commit()
file_record = FileRecord(
filehash="jkl012",
original_filename="french.pdf",
local_filename="/tmp/french.pdf",
file_size=256,
mime_type="application/pdf",
is_duplicate=False,
pipeline_id=explicit_pipeline.id,
)
db_session.add(file_record)
db_session.commit()
result = _get_pipeline_ocr_language(db_session, file_record, owner_id="user1")
assert result == "fra"