from __future__ import annotations import base64 import contextlib import copy import json import os import shutil import subprocess import sys import tempfile import unittest from pathlib import Path from unittest import mock from docforge.errors import DocForgeError from docforge.graph_projection import ( GraphViewRequestV1, build_graph_projection_package, build_graph_view_plan, ) from docforge.manual_projection import ( build_manual_projection_package, build_manual_render_plan, ) from docforge.project import Project from docforge.projection_contract import ( ProjectionPackageV1, ProjectionReceiptV1, canonical_projection_bytes, ) from docforge.projection_worker import ( MAX_WORKER_ARTIFACT_BYTES, MAX_WORKER_RESPONSE_BYTES, _child_response, render_projection_in_worker, ) from docforge.render_contract import GenericHtmlRenderer from docforge_renderers.graph import PortableGraphHtmlRenderer from docforge_renderers.manual import ManualHtmlRenderer ROOT = Path(__file__).resolve().parents[1] FIXTURES = ROOT / "tests" / "fixtures" class ProjectionWorkerTests(unittest.TestCase): def setUp(self) -> None: self.temporary = tempfile.TemporaryDirectory() self.addCleanup(self.temporary.cleanup) self.root = Path(self.temporary.name) / "alpha" shutil.copytree(FIXTURES / "alpha", self.root) self.project = Project.open(self.root) self.snapshot = self.project.load() render = self.snapshot.descriptor.render assert render is not None self.manual_view = render.views[0] self.manual_version = GenericHtmlRenderer().renderer_version self.manual_plan = build_manual_render_plan( self.snapshot, self.manual_view, changeset_hash=None, ) self.manual_package = build_manual_projection_package( self.manual_plan, self.manual_view.template_path.read_bytes(), renderer_id="generic_html", renderer_version=self.manual_version, max_output_bytes=self.snapshot.descriptor.limits.max_render_bytes, ) self.graph_plan = build_graph_view_plan( self.snapshot, GraphViewRequestV1( view_id="architecture", title="Alpha architecture", root_node_id="guide.workflow", depth=2, max_nodes=20, max_edges=40, max_work=1_000, ), False, ) self.graph_package = build_graph_projection_package( self.graph_plan, renderer_id=PortableGraphHtmlRenderer.renderer_id, renderer_version=PortableGraphHtmlRenderer.renderer_version, max_output_bytes=self.snapshot.descriptor.limits.max_render_bytes, ) def test_manual_and_graph_workers_match_the_in_process_renderers(self) -> None: cases = ( ( self.manual_package, ManualHtmlRenderer(self.manual_version).render(self.manual_package), ), ( self.graph_package, PortableGraphHtmlRenderer().render(self.graph_package), ), ) for package, expected in cases: with self.subTest(kind=package.kind): result = render_projection_in_worker(package) self.assertEqual(expected.artifacts, result.artifacts) receipt = result.receipt.as_dict() self.assertEqual(package.package_id, receipt["package_id"]) self.assertEqual(package.document["plan_id"], receipt["plan_id"]) self.assertIs(type(receipt["peak_memory_bytes"]), int) self.assertGreater(receipt["peak_memory_bytes"], 0) def test_parent_launch_is_fixed_and_request_contains_no_runtime_authority(self) -> None: original = subprocess.run with mock.patch( "docforge.projection_worker.subprocess.run", wraps=original, ) as launched: render_projection_in_worker(self.manual_package) command = launched.call_args.args[0] options = launched.call_args.kwargs self.assertEqual( [sys.executable, "-I", "-m", "docforge.projection_worker"], command, ) self.assertIs(options["shell"], False) self.assertEqual(sys.prefix, options["cwd"]) self.assertEqual( {"PYTHONIOENCODING": "utf-8", "PYTHONUTF8": "1"}, options["env"], ) request = options["input"] self.assertEqual( canonical_projection_bytes(self.manual_package.as_dict()) + b"\n", request, ) for forbidden in ( str(self.root).encode(), b"project_path", b"index_path", b"database_path", b'"command"', b'"module"', b'"shell"', b'"sql"', ): self.assertNotIn(forbidden, request) def test_worker_ignores_hostile_cwd_and_pythonpath_import_shadows(self) -> None: shadow_root = Path(self.temporary.name) / "shadow" shadow_package = shadow_root / "docforge" shadow_package.mkdir(parents=True) sentinel = shadow_root / "executed" shadow_package.joinpath("__init__.py").write_text( "from pathlib import Path\n" f"Path({str(sentinel)!r}).write_text('executed', encoding='utf-8')\n", encoding="utf-8", ) with ( contextlib.chdir(shadow_root), mock.patch.dict( os.environ, {"PYTHONPATH": str(shadow_root)}, clear=False, ), ): result = render_projection_in_worker(self.manual_package) self.assertEqual("manual", result.receipt.document["kind"]) self.assertFalse(sentinel.exists()) def test_parent_accepts_only_valid_packages_and_closed_renderer_versions(self) -> None: with self.assertRaises(TypeError): render_projection_in_worker({}) # type: ignore[arg-type] unsupported = ProjectionPackageV1.create( kind="manual", plan=self.manual_plan, renderer={"renderer_id": "generic_html", "renderer_version": "other"}, components=[{"component_id": "manual.document@1"}], assets=copy.deepcopy(self.manual_package.document["assets"]), output_policy=copy.deepcopy(self.manual_package.document["output_policy"]), ) with ( mock.patch("docforge.projection_worker._invoke_worker") as invoke, self.assertRaises(DocForgeError) as raised, ): render_projection_in_worker(unsupported) self.assertEqual("unsupported_renderer", raised.exception.code) invoke.assert_not_called() permissive_allowance = ProjectionPackageV1.create( kind="manual", plan=self.manual_plan, renderer={ "renderer_id": "generic_html", "renderer_version": self.manual_version, }, components=[{"component_id": "manual.document@1"}], assets=copy.deepcopy(self.manual_package.document["assets"]), output_policy={ "artifact_ids": ["manual.html"], "max_total_bytes": MAX_WORKER_ARTIFACT_BYTES + 1, }, ) rendered = render_projection_in_worker(permissive_allowance) self.assertEqual( ManualHtmlRenderer(self.manual_version) .render(permissive_allowance) .artifacts[0] .content, rendered.artifacts[0].content, ) def test_process_timeout_exit_and_signal_fail_closed(self) -> None: failures = ( ( subprocess.TimeoutExpired(["worker"], 30), "projection_worker_timeout", ), ( subprocess.CompletedProcess(["worker"], 2, stdout=b""), "projection_worker_failure", ), ( subprocess.CompletedProcess(["worker"], -9, stdout=b""), "projection_worker_failure", ), ) for outcome, code in failures: with self.subTest(outcome=type(outcome).__name__): patch = ( mock.patch( "docforge.projection_worker._invoke_worker", side_effect=outcome, ) if isinstance(outcome, BaseException) else mock.patch( "docforge.projection_worker._invoke_worker", return_value=outcome, ) ) with patch, self.assertRaises(DocForgeError) as raised: render_projection_in_worker(self.manual_package) self.assertEqual(code, raised.exception.code) def test_malformed_trailing_and_oversized_responses_fail_closed(self) -> None: valid = _child_response(self.manual_package) malformed = ( b"", b"{}", b"not-json\n", b'{ "schema_version": 1 }\n', valid + b"{}\n", b"x" * (MAX_WORKER_RESPONSE_BYTES + 1), ) for response in malformed: with ( self.subTest(length=len(response)), mock.patch( "docforge.projection_worker._invoke_worker", return_value=subprocess.CompletedProcess( ["worker"], 0, stdout=response, ), ), self.assertRaises(DocForgeError) as raised, ): render_projection_in_worker(self.manual_package) self.assertEqual("projection_worker_failure", raised.exception.code) def test_wrong_artifact_and_receipt_evidence_fail_closed(self) -> None: valid = json.loads(_child_response(self.manual_package)) cases: dict[str, dict[str, object]] = {} wrong_content = copy.deepcopy(valid) wrong_content["artifacts"][0]["content_base64"] = base64.b64encode(b"changed").decode() cases["content_hash"] = wrong_content wrong_id = copy.deepcopy(valid) wrong_id["artifacts"][0]["artifact_id"] = "other.html" cases["artifact_id"] = wrong_id missing_artifact = copy.deepcopy(valid) missing_artifact["artifacts"] = [] cases["artifact_count"] = missing_artifact wrong_media_type = copy.deepcopy(valid) wrong_media_type["artifacts"][0]["media_type"] = "application/octet-stream" cases["media_type"] = wrong_media_type bad_base64 = copy.deepcopy(valid) bad_base64["artifacts"][0]["content_base64"] = "***" cases["base64"] = bad_base64 wrong_package = copy.deepcopy(valid) receipt = wrong_package["receipt"] replacement = ProjectionReceiptV1.create( kind="manual", package_id="0" * 64, plan_id=receipt["plan_id"], renderer=receipt["renderer"], artifacts=receipt["artifacts"], diagnostics=receipt["diagnostics"], timing=receipt["timing"], peak_memory_bytes=receipt["peak_memory_bytes"], ) wrong_package["receipt"] = replacement.as_dict() cases["package_id"] = wrong_package zero_peak = copy.deepcopy(valid) receipt = zero_peak["receipt"] replacement = ProjectionReceiptV1.create( kind="manual", package_id=receipt["package_id"], plan_id=receipt["plan_id"], renderer=receipt["renderer"], artifacts=receipt["artifacts"], diagnostics=receipt["diagnostics"], timing=receipt["timing"], peak_memory_bytes=0, ) zero_peak["receipt"] = replacement.as_dict() cases["peak_memory"] = zero_peak wrong_size = copy.deepcopy(valid) receipt = wrong_size["receipt"] artifacts = copy.deepcopy(receipt["artifacts"]) artifacts[0]["bytes"] += 1 replacement = ProjectionReceiptV1.create( kind="manual", package_id=receipt["package_id"], plan_id=receipt["plan_id"], renderer=receipt["renderer"], artifacts=artifacts, diagnostics=receipt["diagnostics"], timing=receipt["timing"], peak_memory_bytes=receipt["peak_memory_bytes"], ) wrong_size["receipt"] = replacement.as_dict() cases["artifact_size"] = wrong_size configured_oversize = copy.deepcopy(valid) maximum = self.manual_package.document["output_policy"]["max_total_bytes"] configured_oversize["artifacts"][0]["content_base64"] = base64.b64encode( b"x" * (maximum + 1) ).decode() cases["configured_aggregate"] = configured_oversize for name, document in cases.items(): response = canonical_projection_bytes(document) + b"\n" with ( self.subTest(name=name), mock.patch( "docforge.projection_worker._invoke_worker", return_value=subprocess.CompletedProcess( ["worker"], 0, stdout=response, ), ), self.assertRaises(DocForgeError) as raised, ): render_projection_in_worker(self.manual_package) self.assertEqual("projection_worker_failure", raised.exception.code) def test_child_rejects_noncanonical_invalid_and_trailing_requests(self) -> None: valid = canonical_projection_bytes(self.manual_package.as_dict()) + b"\n" requests = ( b"{}\n", b'{ "schema_version": 1 }\n', valid + b"{}\n", ) for request in requests: with self.subTest(length=len(request)): completed = subprocess.run( [sys.executable, "-m", "docforge.projection_worker"], cwd=ROOT, input=request, capture_output=True, check=False, timeout=10, ) self.assertEqual(2, completed.returncode) self.assertEqual(b"", completed.stdout) if __name__ == "__main__": unittest.main()