Merge branch 'main' into copilot/add-document-sharing-feature
This commit is contained in:
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user