import os import tempfile import unittest from unittest import mock from scripts import oracle_skill as skill class FakeResponse: def __init__(self, chunks=None, payload=None, ok=True, status_code=200, headers=None): self._chunks = chunks or [] self._payload = payload or {} self.ok = ok self.status_code = status_code self.headers = headers or {} self.text = "" def __enter__(self): return self def __exit__(self, exc_type, exc, tb): return False def json(self): return self._payload def iter_content(self, chunk_size=65536): del chunk_size return iter(self._chunks) class AWRSkillTests(unittest.TestCase): def test_default_transit_url_uses_https_domain(self): self.assertEqual(skill.TRANSIT_URL, "https://ts.henlo.net") def config(self): return { "access_token": "test-token", "expires_at": "2999-01-01T00:00:00+08:00", "transit_url": "http://transit.example", "client_code": "WEIRUI", } def test_awr_status_uses_authenticated_dedicated_endpoint(self): response = FakeResponse(payload={"success": True, "data": {"status": "ready"}}) with mock.patch.object(skill, "get_config", return_value=self.config()), mock.patch.object( skill.requests, "get", return_value=response ) as get: result = skill.awr_status("WEIRUI") self.assertTrue(result["success"]) self.assertEqual(get.call_args.args[0], "http://transit.example/api/awr/status") self.assertEqual(get.call_args.kwargs["params"], {"client_code": "WEIRUI"}) self.assertEqual(get.call_args.kwargs["headers"]["Authorization"], "Bearer test-token") def test_awr_status_returns_schema_resolved_manual_dba_sql_when_permission_is_missing(self): response = FakeResponse( payload={ "success": True, "client_code": "WEIRUI", "data": {"status": "permission_required", "reason": "AWR export function is missing"}, } ) with mock.patch.object(skill, "get_config", return_value=self.config()), mock.patch.object( skill.requests, "get", return_value=response ), mock.patch.object(skill, "query", return_value={"success": True, "data": {"schema": "bosnds3"}}): result = skill.awr_status("WEIRUI") self.assertTrue(result["dba_sql_required"]) self.assertEqual(result["agent_schema"], "BOSNDS3") self.assertIn("GRANT EXECUTE ON SYS.HENLO_AWR_EXPORT TO BOSNDS3", result["dba_sql"]) self.assertIn("GRANT SELECT ON SYS.DBA_HIST_SNAPSHOT TO BOSNDS3", result["dba_sql"]) self.assertIn("IF R.OUTPUT IS NOT NULL THEN", result["dba_sql"]) self.assertIn("DBMS_LOB.WRITEAPPEND(L_HTML, 1, CHR(10))", result["dba_sql"]) self.assertNotIn("", result["dba_sql"]) def test_awr_status_never_emits_placeholder_when_agent_schema_is_unavailable(self): response = FakeResponse( payload={"success": True, "data": {"status": "permission_required", "reason": "permission denied"}} ) with mock.patch.object(skill, "get_config", return_value=self.config()), mock.patch.object( skill.requests, "get", return_value=response ), mock.patch.object(skill, "query", return_value={"success": True, "data": {"schema": ""}}): result = skill.awr_status("WEIRUI") self.assertTrue(result["dba_sql_required"]) self.assertEqual(result["dba_sql"], "") self.assertIn("oracle.schema", result["next_step"]) def test_download_preserves_exact_html_bytes(self): content = "未芮 AWR".encode("utf-8") response = FakeResponse(chunks=[content[:10], content[10:]], headers={"Content-Length": str(len(content))}) with tempfile.TemporaryDirectory() as root: output = os.path.join(root, "awr.html") with mock.patch.object(skill, "get_config", return_value=self.config()), mock.patch.object( skill.requests, "get", return_value=response ) as get: result = skill.download_awr_report("WEIRUI", "20260720", output) self.assertTrue(result["success"], result) with open(output, "rb") as f: self.assertEqual(f.read(), content) self.assertEqual(get.call_args.args[0], "http://transit.example/api/awr/download") self.assertEqual(get.call_args.kwargs["params"], {"client_code": "WEIRUI", "date": "20260720"}) def test_download_rejects_invalid_date_before_network(self): with mock.patch.object(skill, "get_config", return_value=self.config()), mock.patch.object( skill.requests, "get" ) as get: result = skill.download_awr_report("WEIRUI", "../secret") self.assertFalse(result["success"]) get.assert_not_called() def test_download_removes_partial_file_when_size_limit_is_exceeded(self): response = FakeResponse(chunks=[b"123456"]) with tempfile.TemporaryDirectory() as root: output = os.path.join(root, "awr.html") with mock.patch.object(skill, "MAX_AWR_DOWNLOAD_BYTES", 5), mock.patch.object( skill, "get_config", return_value=self.config() ), mock.patch.object(skill.requests, "get", return_value=response): result = skill.download_awr_report("WEIRUI", "20260720", output) self.assertFalse(result["success"]) self.assertFalse(os.path.exists(output)) self.assertFalse(os.path.exists(output + ".part")) if __name__ == "__main__": unittest.main()