conftest.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214
  1. import http.server
  2. import json
  3. import os
  4. import subprocess
  5. import sys
  6. import threading
  7. import time
  8. from pathlib import Path
  9. import pytest
  10. HOOKS_DIR = Path(__file__).resolve().parent.parent / "hooks"
  11. HOOK_SCRIPT = HOOKS_DIR / "security_reminder_hook.py"
  12. sys.path.insert(0, str(HOOKS_DIR))
  13. GIT_ENV = {
  14. "GIT_AUTHOR_NAME": "t", "GIT_AUTHOR_EMAIL": "t@example.com",
  15. "GIT_COMMITTER_NAME": "t", "GIT_COMMITTER_EMAIL": "t@example.com",
  16. "GIT_CONFIG_GLOBAL": os.devnull, "GIT_CONFIG_SYSTEM": os.devnull,
  17. }
  18. def git(cwd, *args):
  19. r = subprocess.run(
  20. ["git", *args], cwd=cwd, capture_output=True, text=True,
  21. env={**os.environ, **GIT_ENV},
  22. )
  23. assert r.returncode == 0, r.stderr
  24. return r.stdout
  25. def make_repo(path, files=None):
  26. path.mkdir(parents=True, exist_ok=True)
  27. git(path, "init", "-q", "-b", "main")
  28. git(path, "config", "commit.gpgsign", "false")
  29. for name, content in (files or {"README.md": "init\n"}).items():
  30. p = path / name
  31. p.parent.mkdir(parents=True, exist_ok=True)
  32. p.write_text(content)
  33. git(path, "add", "-A")
  34. git(path, "commit", "-q", "-m", "init")
  35. return path
  36. def commit_file(repo, name, content, msg="change"):
  37. p = repo / name
  38. p.parent.mkdir(parents=True, exist_ok=True)
  39. p.write_text(content)
  40. git(repo, "add", "-A")
  41. out = subprocess.run(
  42. ["git", "commit", "-m", msg], cwd=repo, capture_output=True, text=True,
  43. env={**os.environ, **GIT_ENV},
  44. )
  45. assert out.returncode == 0, out.stderr
  46. sha = git(repo, "rev-parse", "HEAD").strip()
  47. return sha, out.stdout + out.stderr
  48. VULN_PY = (
  49. "import subprocess\n"
  50. "def run(user):\n"
  51. " subprocess.call('ls ' + user, shell=True)\n"
  52. )
  53. @pytest.fixture
  54. def workspace(tmp_path):
  55. ws = tmp_path / "ws"
  56. ws.mkdir()
  57. repo = make_repo(ws / "sub", {"app.py": "print('hi')\n"})
  58. (ws / "node_modules" / "junk").mkdir(parents=True)
  59. return ws, repo
  60. class _Stub(http.server.BaseHTTPRequestHandler):
  61. def do_POST(self):
  62. n = int(self.headers.get("Content-Length") or 0)
  63. if n:
  64. self.rfile.read(n)
  65. self.server.calls.append(self.path)
  66. if self.server.delay:
  67. time.sleep(self.server.delay)
  68. if self.server.status == 200:
  69. vulns = list(self.server.vulns)
  70. body = json.dumps({
  71. "id": "msg_stub", "type": "message", "role": "assistant",
  72. "model": "stub", "stop_reason": "end_turn",
  73. "content": [{"type": "text", "text": json.dumps(
  74. {"hasVulnerabilities": bool(vulns), "vulnerabilities": vulns})}],
  75. "usage": {"input_tokens": 1, "output_tokens": 1},
  76. }).encode()
  77. else:
  78. body = b'{"type":"error","error":{"type":"invalid_request_error","message":"stub"}}'
  79. self.send_response(self.server.status)
  80. self.send_header("Content-Type", "application/json")
  81. self.send_header("Content-Length", str(len(body)))
  82. self.end_headers()
  83. self.wfile.write(body)
  84. do_HEAD = do_GET = do_POST
  85. def log_message(self, *a):
  86. pass
  87. @pytest.fixture
  88. def stub_api():
  89. srv = http.server.HTTPServer(("127.0.0.1", 0), _Stub)
  90. srv.calls = []
  91. srv.status = 200
  92. srv.delay = 0
  93. srv.vulns = []
  94. t = threading.Thread(target=srv.serve_forever, daemon=True)
  95. t.start()
  96. try:
  97. yield srv
  98. finally:
  99. srv.shutdown()
  100. @pytest.fixture
  101. def hook_env(tmp_path, stub_api):
  102. state = tmp_path / "state"
  103. state.mkdir()
  104. env = {k: v for k, v in os.environ.items()
  105. if k not in ("ANTHROPIC_AUTH_TOKEN", "CLAUDE_CODE_REMOTE",
  106. "CLAUDE_PROJECT_DIR", "HTTP_PROXY", "HTTPS_PROXY",
  107. "http_proxy", "https_proxy", "ALL_PROXY", "all_proxy",
  108. "CLAUDE_CODE_USE_BEDROCK", "CLAUDE_CODE_USE_VERTEX",
  109. "CLAUDE_CODE_USE_FOUNDRY")}
  110. env.update(GIT_ENV)
  111. env.update({
  112. "SECURITY_WARNINGS_STATE_DIR": str(state),
  113. "ANTHROPIC_API_KEY": "test-key",
  114. "ANTHROPIC_BASE_URL": f"http://127.0.0.1:{stub_api.server_port}",
  115. "NO_PROXY": "*", "no_proxy": "*",
  116. "SG_AGENTIC_COMMIT_REVIEW": "0",
  117. "SECURITY_GUIDANCE_COMMIT_REVIEW": "on",
  118. "SG_PUSH_SWEEP": "on",
  119. "PYTHONDONTWRITEBYTECODE": "1",
  120. })
  121. return env
  122. def run_hook(payload, env, python=sys.executable):
  123. r = subprocess.run(
  124. [python, str(HOOK_SCRIPT)], input=json.dumps(payload),
  125. capture_output=True, text=True, env=env, timeout=120,
  126. )
  127. return r.returncode, r.stdout, r.stderr
  128. STUB_VULN = {
  129. "filePath": "app.py", "category": "command_injection", "severity": "high",
  130. "vulnerableCode": "subprocess.call('ls ' + user, shell=True)",
  131. "description": "user input reaches a shell", "recommendation": "pass an argv list",
  132. }
  133. def metrics_of(stdout):
  134. for line in stdout.splitlines():
  135. line = line.strip()
  136. if line.startswith("{"):
  137. try:
  138. m = json.loads(line).get("metrics")
  139. except json.JSONDecodeError:
  140. continue
  141. if m is not None:
  142. return m
  143. return None
  144. def bash_payload(cwd, command, stdout="", stderr="", session_id="s1", tool_use_id=None):
  145. p = {
  146. "session_id": session_id,
  147. "hook_event_name": "PostToolUse",
  148. "tool_name": "Bash",
  149. "tool_input": {"command": command},
  150. "tool_response": {"stdout": stdout, "stderr": stderr, "interrupted": False},
  151. "cwd": str(cwd),
  152. }
  153. if tool_use_id:
  154. p["tool_use_id"] = tool_use_id
  155. return p
  156. def edit_payload(cwd, file_path, new_string="x", session_id="s1"):
  157. return {
  158. "session_id": session_id,
  159. "hook_event_name": "PostToolUse",
  160. "tool_name": "Edit",
  161. "tool_input": {"file_path": str(file_path), "old_string": "", "new_string": new_string},
  162. "tool_response": {},
  163. "cwd": str(cwd),
  164. }
  165. def stop_payload(cwd, event="Stop", session_id="s1"):
  166. return {
  167. "session_id": session_id,
  168. "hook_event_name": event,
  169. "stop_hook_active": False,
  170. "cwd": str(cwd),
  171. }
  172. def ups_payload(cwd, session_id="s1"):
  173. return {
  174. "session_id": session_id,
  175. "hook_event_name": "UserPromptSubmit",
  176. "prompt": "hi",
  177. "cwd": str(cwd),
  178. }