mirror of
https://github.com/usestrix/strix.git
synced 2026-10-01 02:03:55 +00:00
269 lines
9.7 KiB
Python
269 lines
9.7 KiB
Python
"""Real Strix loop + SDK shell/filesystem + customer tests, with scripted inference."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import shlex
|
|
import sys
|
|
import zipfile
|
|
from types import SimpleNamespace
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
import pytest
|
|
from agents import Agent, Model, RunConfig, RunContextWrapper
|
|
from agents.exceptions import MaxTurnsExceeded
|
|
from agents.items import ModelResponse
|
|
from agents.sandbox import SandboxRunConfig
|
|
from agents.tool import CustomTool
|
|
from agents.usage import Usage
|
|
from openai.types.responses import (
|
|
ResponseCustomToolCall,
|
|
ResponseFunctionToolCall,
|
|
ResponseOutputMessage,
|
|
ResponseOutputText,
|
|
)
|
|
|
|
import strix.core.hooks as hooks_module
|
|
from strix.config.models import _completed_stream_event
|
|
from strix.core.hooks import BudgetExceededError, ReportUsageHooks
|
|
from strix.fix import PreparationState
|
|
from strix.fix import runtime as fix_runtime
|
|
from strix.fix.runtime import _FixHooks
|
|
from strix.interface.fix_cli import _summary
|
|
from tests.test_fix_reliability import environment, existing_suite
|
|
from tests.test_fix_runtime import _request, _workspace
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
from pathlib import Path
|
|
|
|
|
|
def call(name: str, **arguments: Any) -> ResponseFunctionToolCall:
|
|
return ResponseFunctionToolCall(
|
|
type="function_call", name=name, call_id=name, arguments=json.dumps(arguments)
|
|
)
|
|
|
|
|
|
def finish(outcome: str, summary: str = "Fix and validation results reviewed.") -> Any:
|
|
return call("agent_finish", outcome=outcome, result_summary=summary)
|
|
|
|
|
|
def shell(cmd: str) -> Any:
|
|
return call("exec_command", cmd=cmd, login=False, yield_time_ms=10000)
|
|
|
|
|
|
def patch(value: str = "safe") -> list[Any]:
|
|
production = "def result():\n return " + repr(value) + "\n"
|
|
regression = (
|
|
"import unittest\nfrom app import result\nclass Security(unittest.TestCase):\n"
|
|
" def test_safe(self): self.assertEqual(result(),'safe')\n"
|
|
)
|
|
return [
|
|
shell(f"printf %s {shlex.quote(production)} > app.py"),
|
|
shell(f"printf %s {shlex.quote(regression)} > tests/test_security.py"),
|
|
]
|
|
|
|
|
|
def suite_commands() -> list[Any]:
|
|
python = shlex.quote(sys.executable)
|
|
return [
|
|
shell(f"{python} -m unittest discover -s tests -p test_existing.py"),
|
|
shell(f"{python} -m unittest discover -s tests -p test_security.py"),
|
|
]
|
|
|
|
|
|
class ScriptedModel(Model):
|
|
def __init__(self, repair: list[Any], review: list[Any] | None = None) -> None:
|
|
self.responses = {"repair": repair, "review": review or []}
|
|
self.inputs: dict[str, list[Any]] = {"repair": [], "review": []}
|
|
self.tools: set[str] = set()
|
|
|
|
async def get_response(self, **kwargs: Any) -> ModelResponse:
|
|
role = "review" if "Independently review" in kwargs["system_instructions"] else "repair"
|
|
self.inputs[role].append(list(kwargs["input"]))
|
|
self.tools.update(t.name for t in kwargs["tools"])
|
|
assert self.responses[role], f"Unexpected additional {role} turn"
|
|
item = self.responses[role].pop(0)
|
|
if isinstance(item, str):
|
|
item = ResponseOutputMessage(
|
|
id=f"msg-{role}-{len(self.inputs[role])}",
|
|
type="message",
|
|
role="assistant",
|
|
status="completed",
|
|
content=[ResponseOutputText(type="output_text", text=item, annotations=[])],
|
|
)
|
|
else:
|
|
arguments = json.loads(item.arguments)
|
|
item = item.model_copy(
|
|
update={
|
|
"call_id": f"{role}-{len(self.inputs[role])}",
|
|
"arguments": json.dumps(arguments),
|
|
}
|
|
)
|
|
if item.name == "apply_patch" and any(
|
|
isinstance(t, CustomTool) and t.name == item.name for t in kwargs["tools"]
|
|
):
|
|
item = ResponseCustomToolCall(
|
|
type="custom_tool_call",
|
|
name=item.name,
|
|
call_id=item.call_id,
|
|
input=arguments["patch"],
|
|
)
|
|
return ModelResponse(output=[item], usage=Usage(requests=1), response_id=None)
|
|
|
|
async def stream_response(self, *args: Any, **kwargs: Any) -> Any:
|
|
kwargs.update(
|
|
zip(
|
|
[
|
|
"system_instructions",
|
|
"input",
|
|
"model_settings",
|
|
"tools",
|
|
"output_schema",
|
|
"handoffs",
|
|
"tracing",
|
|
],
|
|
args,
|
|
strict=False,
|
|
)
|
|
)
|
|
yield _completed_stream_event(await self.get_response(**kwargs), "scripted")
|
|
|
|
|
|
async def scenario(
|
|
tmp_path: Path,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
model: ScriptedModel,
|
|
turns: int = 30,
|
|
*,
|
|
repair_turns: int | None = None,
|
|
review_turns: int | None = None,
|
|
) -> tuple[Any, Any]:
|
|
workspace, _ = _workspace(tmp_path)
|
|
commit = existing_suite(workspace)
|
|
env = environment(workspace, tmp_path)
|
|
monkeypatch.setattr(
|
|
fix_runtime,
|
|
"_run_config",
|
|
lambda env: RunConfig(
|
|
model=model, sandbox=SandboxRunConfig(session=env.session), tracing_disabled=True
|
|
),
|
|
)
|
|
request = _request(commit)
|
|
request.max_agent_turns = turns
|
|
request.max_repair_turns = repair_turns
|
|
request.max_review_turns = review_turns
|
|
result = await fix_runtime.run_fix_preparation(
|
|
request,
|
|
workspace,
|
|
sandbox_session=env.session,
|
|
runtime_environment=env,
|
|
artifact_path=tmp_path / "prepared.zip",
|
|
)
|
|
return result, env
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_single_agent_implements_runs_both_test_suites_and_exports(tmp_path, monkeypatch):
|
|
model = ScriptedModel([*patch(), *suite_commands(), finish("done", "Both suites passed")])
|
|
result, _env = await scenario(tmp_path, monkeypatch, model)
|
|
assert result.state is PreparationState.READY, result.model_dump_json()
|
|
assert result.validation_mode == "single_agent"
|
|
assert not model.inputs["review"]
|
|
assert result.verifier is None
|
|
assert result.completion.turns_used == 5
|
|
assert all("Ran 1 test" in c.output for c in result.checks[-2:])
|
|
assert {"exec_command", "apply_patch", "agent_finish"} <= model.tools
|
|
assert not {"create_agent", "finish_scan", "record_coverage"} & model.tools
|
|
with zipfile.ZipFile(tmp_path / "prepared.zip") as artifact:
|
|
assert "files/tests/test_security.py" in artifact.namelist()
|
|
assert "Both suites passed" in _summary(result)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_agent_corrects_failed_test_in_same_conversation(tmp_path, monkeypatch):
|
|
model = ScriptedModel(
|
|
[*patch("still unsafe"), *suite_commands(), *patch(), *suite_commands(), finish("done")]
|
|
)
|
|
result, _ = await scenario(tmp_path, monkeypatch, model)
|
|
assert result.state is PreparationState.READY
|
|
assert any(c.exit_code != 0 for c in result.checks)
|
|
assert all(c.exit_code == 0 for c in result.checks[-2:])
|
|
assert not model.inputs["review"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("end", ["blocked", "limit"])
|
|
async def test_blocked_or_capped_agent_discards_patch(tmp_path, monkeypatch, end):
|
|
model = ScriptedModel([*patch(), finish("blocked", "Database unavailable")])
|
|
result, _ = await scenario(tmp_path, monkeypatch, model, turns=2 if end == "limit" else 10)
|
|
assert result.state is PreparationState.BLOCKED
|
|
assert not result.final_file_manifest
|
|
assert not (tmp_path / "prepared.zip").exists()
|
|
assert len(model.inputs["repair"]) <= (2 if end == "limit" else 3)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_native_lifecycle_retries_invalid_outcome(tmp_path, monkeypatch):
|
|
model = ScriptedModel([*patch(), finish("approved"), *suite_commands(), finish("done")])
|
|
result, _ = await scenario(tmp_path, monkeypatch, model)
|
|
assert result.state is PreparationState.READY
|
|
assert result.completion.turns_used == 6
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resume_consumes_remaining_turn_allowance(tmp_path):
|
|
|
|
env = environment(tmp_path / "source", tmp_path)
|
|
env.turns_used = 299
|
|
env.max_repair_turns = 500
|
|
saved = []
|
|
env.turn_sink = saved.append
|
|
hooks = _FixHooks(env)
|
|
|
|
context = RunContextWrapper(context={})
|
|
await hooks.on_llm_start(context, Agent(name="fix"), "", [])
|
|
with pytest.raises(MaxTurnsExceeded):
|
|
await hooks.on_llm_start(context, Agent(name="fix"), "", [])
|
|
assert saved == [300]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_completion_recommendations_are_in_summary(tmp_path, monkeypatch):
|
|
model = ScriptedModel(
|
|
[
|
|
*patch(),
|
|
*suite_commands(),
|
|
call(
|
|
"agent_finish",
|
|
outcome="done",
|
|
result_summary="Fixed and tested",
|
|
final_recommendations=["Run the nightly suite"],
|
|
),
|
|
]
|
|
)
|
|
result, _ = await scenario(tmp_path, monkeypatch, model)
|
|
assert "Run the nightly suite" in _summary(result)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_fix_respects_live_scan_budget_without_double_counting(tmp_path, monkeypatch):
|
|
|
|
recorded = []
|
|
state = SimpleNamespace(
|
|
get_total_llm_cost=lambda: 2.0, record_sdk_usage=lambda **kwargs: recorded.append(kwargs)
|
|
)
|
|
monkeypatch.setattr(hooks_module, "get_global_report_state", lambda: state)
|
|
env = environment(tmp_path / "source", tmp_path)
|
|
shared = ReportUsageHooks(model="test", max_budget_usd=10)
|
|
env.scan_hooks = shared
|
|
hooks = _FixHooks(env)
|
|
context = RunContextWrapper(context={"agent_id": "fix", "parent_id": "root"})
|
|
shared.set_max_budget_usd(1)
|
|
with pytest.raises(BudgetExceededError):
|
|
await hooks.on_llm_end(
|
|
context,
|
|
Agent(name="fix"),
|
|
ModelResponse(output=[], usage=Usage(requests=1), response_id=None),
|
|
)
|
|
assert len(recorded) == 1
|