mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 17:27:26 +02:00
Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
597aae6715 | ||
|
|
06b158d1fa | ||
|
|
c29eb73c7f | ||
|
|
72833b8e43 | ||
|
|
1117ba6d4a | ||
|
|
53e4658d88 | ||
|
|
58df71d3db | ||
|
|
0b9e029a5d | ||
|
|
b260a4ee38 | ||
|
|
f8a8801d56 |
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "strix-agent"
|
||||
version = "1.5.1"
|
||||
version = "1.5.2"
|
||||
description = "Open-source AI Hackers for your apps"
|
||||
readme = "README.md"
|
||||
license = "Apache-2.0"
|
||||
|
||||
@@ -160,7 +160,9 @@ def _schema_types(spec: dict[str, Any]) -> set[str]:
|
||||
def _decode_structured(value: str, types: set[str]) -> Any:
|
||||
stripped = value.strip()
|
||||
if not stripped:
|
||||
return value
|
||||
# An empty string is the model's "no value" for a list/dict param; give it
|
||||
# the empty container so it validates instead of failing the type check.
|
||||
return [] if "array" in types else {}
|
||||
try:
|
||||
decoded = json.loads(stripped)
|
||||
except json.JSONDecodeError:
|
||||
|
||||
@@ -39,6 +39,8 @@ INTERACTIVE BEHAVIOR:
|
||||
- To end the whole engagement, call the lifecycle tool: finish_scan (root) or agent_finish (subagent).
|
||||
- A turn that ends with plain text and no tool call does NOT stop you: the system nudges you to continue and will re-run you. Do not rely on going silent to pause — it will not pause you.
|
||||
- Answering a user question: put the answer in respond_to_user's message. Do not write the answer as plain text and then fall silent — that does not reach a stopping point, it just triggers a continuation nudge.
|
||||
- If all you want to do is reply and stop, that whole turn is ONE respond_to_user call carrying the answer. Do not write the answer as text and then call respond_to_user as well: the user reads it twice.
|
||||
- If you do end a turn on plain text and the nudge arrives, your words already reached the user. Do not restate them: call respond_to_user with NO message to simply wait, or with only whatever you still need to add.
|
||||
- You may include brief explanatory text before a tool call, and you can narrate while you work — plain text is shown to the user as you go. Narrating is free; respond_to_user is specifically the act of WAITING for the user, so do not call it just to give a status update.
|
||||
- Respond naturally when the user asks questions or gives instructions.
|
||||
- While actively working on a task, every turn should carry exactly one tool call — use think to plan, the appropriate tool to act, and respond_to_user only when you genuinely need the user.
|
||||
|
||||
@@ -830,7 +830,7 @@ async def _append_tool_required_message(
|
||||
"execution and never hands control to the user: it is shown to the user, and the "
|
||||
"run continues. Continue immediately and call exactly one tool. "
|
||||
"If you have something to tell the user and nothing to do until they reply, "
|
||||
"call respond_to_user. "
|
||||
"call respond_to_user — with no message if you have already said it. "
|
||||
"If you are blocked waiting for another agent, call wait_for_agents. "
|
||||
f"If the whole engagement is complete, call {finish_tool}. "
|
||||
"Otherwise use the appropriate execution or planning tool. "
|
||||
|
||||
@@ -138,6 +138,15 @@ def build_root_task(scan_config: dict[str, Any]) -> str:
|
||||
"target to assess: the instructions below are the only source of "
|
||||
"truth for what to do."
|
||||
)
|
||||
elif not parts and user_instructions:
|
||||
# Neither a target nor a directory, but there is an instruction: the user
|
||||
# declined the mount, so the instruction is all there is. Say so, or the
|
||||
# agent goes looking for a scope that was never given.
|
||||
parts.append(
|
||||
"\n\nNo scan target and no working directory were provided. The "
|
||||
"instructions below are the only source of truth for what to do; "
|
||||
"work from them and from what you can reach yourself."
|
||||
)
|
||||
|
||||
parts.extend(_render_diff_scope(diff_scope))
|
||||
|
||||
|
||||
+12
-3
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
@@ -429,7 +430,6 @@ async def run_strix_scan(
|
||||
except BudgetExceededError as exc:
|
||||
logger.info("Scan %s stopped: %s", scan_id, exc)
|
||||
if root_id is not None:
|
||||
await coordinator.cancel_descendants(root_id)
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator.set_status(root_id, "stopped")
|
||||
return None
|
||||
@@ -442,19 +442,28 @@ async def run_strix_scan(
|
||||
scan_id,
|
||||
)
|
||||
if root_id is not None:
|
||||
await coordinator.cancel_descendants(root_id)
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator.set_status(root_id, "stopped")
|
||||
return None
|
||||
except (asyncio.CancelledError, KeyboardInterrupt):
|
||||
logger.info("Scan %s interrupted by the user", scan_id)
|
||||
if root_id is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator.set_status(root_id, "running")
|
||||
raise
|
||||
except BaseException:
|
||||
logger.exception("Strix scan %s failed", scan_id)
|
||||
if root_id is not None:
|
||||
await coordinator.cancel_descendants(root_id)
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator.set_status(root_id, "failed")
|
||||
raise
|
||||
finally:
|
||||
configure_spill_writer(None)
|
||||
# Settle descendants before closing sessions: on a clean finish a child
|
||||
# can still be mid-turn, and closing its session underneath it crashes it.
|
||||
if root_id is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator.cancel_descendants(root_id)
|
||||
for s in sessions_to_close:
|
||||
with contextlib.suppress(Exception):
|
||||
s.close()
|
||||
|
||||
+20
-2
@@ -4,6 +4,8 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import sqlite3
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from weakref import WeakKeyDictionary
|
||||
|
||||
@@ -12,7 +14,7 @@ from agents.memory import SQLiteSession
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterator
|
||||
from pathlib import Path
|
||||
|
||||
from agents.items import TResponseInputItem
|
||||
@@ -22,9 +24,25 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _PooledConnectionSession(SQLiteSession):
|
||||
@contextmanager
|
||||
def _locked_connection(self) -> Iterator[sqlite3.Connection]:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise RuntimeError("SQLiteSession is closed")
|
||||
if self._is_memory_db:
|
||||
yield self._shared_connection
|
||||
return
|
||||
connection = sqlite3.connect(str(self.db_path), check_same_thread=False)
|
||||
try:
|
||||
yield connection
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
def open_agent_session(agent_id: str, path: Path) -> SQLiteSession:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
return SQLiteSession(session_id=agent_id, db_path=path)
|
||||
return _PooledConnectionSession(session_id=agent_id, db_path=path)
|
||||
|
||||
|
||||
async def seed_initial_input(session: Session, initial_input: Any) -> bool:
|
||||
|
||||
@@ -328,10 +328,11 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser
|
||||
parser.error(f"--resume {args.resume}: run.json unreadable: {exc}")
|
||||
|
||||
args.targets_info = state.get("targets_info") or []
|
||||
# A target-less run has no targets_info at all: it works in a mounted
|
||||
# directory, driven by its instruction.
|
||||
# A target-less run has no targets_info at all. It is driven by its
|
||||
# instruction, over a mounted working directory or over nothing when the
|
||||
# mount was declined, so either of those is enough to resume it.
|
||||
workspace_mount = state.get("workspace_mount") or None
|
||||
if not args.targets_info and not workspace_mount:
|
||||
if not args.targets_info and not workspace_mount and not state.get("user_instruction"):
|
||||
parser.error(f"--resume {args.resume}: run.json has no targets_info")
|
||||
|
||||
for target in args.targets_info:
|
||||
|
||||
@@ -138,13 +138,6 @@ class TuiController:
|
||||
self.error = detail
|
||||
self.notify_changed()
|
||||
|
||||
def enter_setup(self) -> None:
|
||||
"""Return a session to the start screen, e.g. on a declined mount."""
|
||||
self.setup_mode = True
|
||||
self.scan_started = False
|
||||
self.scan_state = "setup"
|
||||
self.notify_changed()
|
||||
|
||||
def add_message(self, text: str, level: str = "info") -> None:
|
||||
self._append_message(text, level)
|
||||
self.notify_changed()
|
||||
@@ -356,14 +349,12 @@ class TuiController:
|
||||
if not isinstance(approved, bool):
|
||||
raise TypeError("approved must be a boolean")
|
||||
self.pending_workspace_mount = None
|
||||
if not approved:
|
||||
# Nothing was prepared, so return to the start screen untouched.
|
||||
self.workspace_mount = None
|
||||
self.enter_setup()
|
||||
return {"approved": False}
|
||||
self.workspace_mount = mount
|
||||
# Declining skips the mount, it does not abandon the scan. The prompt is
|
||||
# the whole of the input either way; the working directory is only an
|
||||
# extra the agent may look at, so the run goes ahead without one.
|
||||
self.workspace_mount = mount if approved else None
|
||||
await self._begin_scan(self._pending_verify)
|
||||
return {"approved": True}
|
||||
return {"approved": approved}
|
||||
|
||||
async def _send_message(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
agent_id = self._required_string(payload, "agent_id")
|
||||
|
||||
@@ -66,13 +66,10 @@ func (m *Model) submitSetupPrompt(value string) (tea.Model, tea.Cmd) {
|
||||
}
|
||||
|
||||
// answerMountConfirmation replies to the working-directory mount the backend is
|
||||
// waiting on. Declining returns to the start screen, so the prompt goes back in
|
||||
// the composer to be edited or given a target instead.
|
||||
// waiting on. Either answer starts the scan - declining only means it runs
|
||||
// without the directory - so the prompt stays with the run rather than coming
|
||||
// back to the composer.
|
||||
func (m *Model) answerMountConfirmation(approved bool) tea.Cmd {
|
||||
if !approved && m.pendingPrompt != "" {
|
||||
m.input.SetValue(m.pendingPrompt)
|
||||
m.resizeViewport()
|
||||
}
|
||||
m.pendingPrompt = ""
|
||||
return send(m.client, "setup.confirm_mount", map[string]any{"approved": approved})
|
||||
}
|
||||
|
||||
@@ -247,13 +247,10 @@ func TestMountConfirmationAnswers(t *testing.T) {
|
||||
if payload.Approved != tc.approved {
|
||||
t.Fatalf("%s: approved=%v, want %v", tc.name, payload.Approved, tc.approved)
|
||||
}
|
||||
// Declining returns to the start screen, so the prompt comes back.
|
||||
want := ""
|
||||
if !tc.approved {
|
||||
want = "find auth bugs in the login flow"
|
||||
}
|
||||
if got := model.input.Value(); got != want {
|
||||
t.Fatalf("%s: composer = %q, want %q", tc.name, got, want)
|
||||
// Either answer launches, so the prompt stays with the run rather than
|
||||
// coming back to the composer.
|
||||
if got := model.input.Value(); got != "" {
|
||||
t.Fatalf("%s: composer = %q, want it cleared", tc.name, got)
|
||||
}
|
||||
if model.pendingPrompt != "" {
|
||||
t.Fatalf("%s: held prompt was not cleared: %q", tc.name, model.pendingPrompt)
|
||||
@@ -290,3 +287,101 @@ func TestSetupPromptWithTargetLaunches(t *testing.T) {
|
||||
t.Fatalf("setup.start (%d) must come after setup.set_instruction (%d): %v", start, instr, types)
|
||||
}
|
||||
}
|
||||
|
||||
// The prompt's buttons are buttons: clicking Cancel has to answer the backend,
|
||||
// which it could not do while the mouse handler had no case for this modal.
|
||||
func TestMountPromptButtonsAreClickable(t *testing.T) {
|
||||
for _, testCase := range []struct {
|
||||
label string
|
||||
approved bool
|
||||
}{
|
||||
{mountConfirmLabel, true},
|
||||
{mountCancelLabel, false},
|
||||
} {
|
||||
connection := &recordingConn{}
|
||||
model := New(&Client{conn: connection})
|
||||
model.width, model.height = 130, 40
|
||||
model.snapshot = protocol.Snapshot{SetupMode: true, WorkingDir: "/Users/me/code/api"}
|
||||
updated, _ := model.submit("find auth bugs in the login flow")
|
||||
model = updated.(Model)
|
||||
connection.Reset()
|
||||
model.snapshot = protocol.Snapshot{
|
||||
ScanStarted: true, ScanState: "preparing", PendingMount: "/Users/me/code/api",
|
||||
}
|
||||
model.syncMountPrompt()
|
||||
|
||||
left, top, panel := model.mountPromptBounds()
|
||||
clicked := false
|
||||
for row, line := range strings.Split(panel, "\n") {
|
||||
plain := ansi.Strip(line)
|
||||
index := strings.Index(plain, testCase.label)
|
||||
if index < 0 {
|
||||
continue
|
||||
}
|
||||
updated, cmd := model.updateModalMouse(tea.MouseMsg{
|
||||
X: left + ansi.StringWidth(plain[:index]) + 1, Y: top + row,
|
||||
Button: tea.MouseButtonLeft, Action: tea.MouseActionPress,
|
||||
})
|
||||
model = updated.(Model)
|
||||
envelopes := drainCommands(t, cmd, connection)
|
||||
if len(envelopes) != 1 || envelopes[0].Type != "setup.confirm_mount" {
|
||||
t.Fatalf("clicking %s sent %v", testCase.label, commandTypes(envelopes))
|
||||
}
|
||||
var payload struct {
|
||||
Approved bool `json:"approved"`
|
||||
}
|
||||
if err := json.Unmarshal(envelopes[0].Payload, &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if payload.Approved != testCase.approved {
|
||||
t.Fatalf("clicking %s answered approved=%v", testCase.label, payload.Approved)
|
||||
}
|
||||
clicked = true
|
||||
break
|
||||
}
|
||||
if !clicked {
|
||||
t.Fatalf("%s was not found in the prompt", testCase.label)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Skipping the mount runs the scan without a directory. It must not throw the
|
||||
// session back to the start screen, and it must not hand the prompt back: the
|
||||
// run has it.
|
||||
func TestSkippingTheMountKeepsTheScanRunning(t *testing.T) {
|
||||
connection := &recordingConn{}
|
||||
model := New(&Client{conn: connection})
|
||||
model.width, model.height = 130, 40
|
||||
model.snapshot = protocol.Snapshot{SetupMode: true, WorkingDir: "/Users/me/code/api"}
|
||||
updated, _ := model.submit("find auth bugs in the login flow")
|
||||
model = updated.(Model)
|
||||
model.snapshot = protocol.Snapshot{
|
||||
ScanStarted: true, ScanState: "preparing", PendingMount: "/Users/me/code/api",
|
||||
}
|
||||
model.syncMountPrompt()
|
||||
if model.modal != modalConfirmMount {
|
||||
t.Fatal("the prompt did not open")
|
||||
}
|
||||
|
||||
model.modalChoice = 1
|
||||
updated, _ = model.updateModal(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
model = updated.(Model)
|
||||
|
||||
// The backend answers by starting the scan with no mount.
|
||||
model.handleEnvelope(stateEnvelope(t, 2, protocol.Snapshot{
|
||||
ScanStarted: true, ScanState: "running",
|
||||
}))
|
||||
|
||||
if model.modal != modalNone {
|
||||
t.Fatalf("the prompt is still open: %v", model.modal)
|
||||
}
|
||||
if model.snapshot.SetupMode {
|
||||
t.Fatal("skipping the mount fell back to the start screen")
|
||||
}
|
||||
if got := model.input.Value(); got != "" {
|
||||
t.Fatalf("the prompt came back to the composer: %q", got)
|
||||
}
|
||||
if model.pendingPrompt != "" {
|
||||
t.Fatalf("the held prompt was not released: %q", model.pendingPrompt)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -474,6 +474,18 @@ func (m Model) updateModalMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) {
|
||||
m.modalChoice = 1
|
||||
return m.updateModal(tea.KeyMsg{Type: tea.KeyEnter})
|
||||
}
|
||||
case modalConfirmMount:
|
||||
left, top, panel := m.mountPromptBounds()
|
||||
if labelHitAt(panel, mountConfirmLabel, left, top, msg.X, msg.Y) {
|
||||
m.modalChoice = 0
|
||||
cmd := m.answerMountConfirmation(true)
|
||||
return m, cmd
|
||||
}
|
||||
if labelHitAt(panel, mountCancelLabel, left, top, msg.X, msg.Y) {
|
||||
m.modalChoice = 1
|
||||
cmd := m.answerMountConfirmation(false)
|
||||
return m, cmd
|
||||
}
|
||||
case modalVulnerability:
|
||||
for _, button := range m.reportButtons() {
|
||||
if button == reportCopy || button == reportDone {
|
||||
@@ -507,7 +519,14 @@ func (m Model) centeredViewBounds(view string) (left, top, width, height int) {
|
||||
|
||||
func (m Model) centeredLabelHit(view, label string, x, y int) bool {
|
||||
left, top, _, _ := m.centeredViewBounds(view)
|
||||
for row, line := range strings.Split(view, "\n") {
|
||||
return labelHitAt(view, label, left, top, x, y)
|
||||
}
|
||||
|
||||
// labelHitAt reports whether a click landed on a label drawn in a panel whose
|
||||
// top-left corner is at (left, top). The mount prompt is docked in a corner
|
||||
// rather than centered, so it cannot use the centered bounds.
|
||||
func labelHitAt(panel, label string, left, top, x, y int) bool {
|
||||
for row, line := range strings.Split(panel, "\n") {
|
||||
plain := ansi.Strip(line)
|
||||
index := strings.Index(plain, label)
|
||||
if index < 0 || y != top+row {
|
||||
@@ -593,7 +612,8 @@ func (m Model) updateModal(key tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
case "esc":
|
||||
if m.modal == modalConfirmMount {
|
||||
// The backend is waiting on an answer; escape declines it.
|
||||
return m, m.answerMountConfirmation(false)
|
||||
cmd := m.answerMountConfirmation(false)
|
||||
return m, cmd
|
||||
}
|
||||
m.closeModal()
|
||||
return m, nil
|
||||
@@ -604,7 +624,10 @@ func (m Model) updateModal(key tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
modal, choice := m.modal, m.modalChoice
|
||||
if modal == modalConfirmMount {
|
||||
// The snapshot closes this prompt once the backend has the answer.
|
||||
return m, m.answerMountConfirmation(choice == 0)
|
||||
// Bound to a variable first: the call restores the held prompt into
|
||||
// the composer, and that has to be in the model being returned.
|
||||
cmd := m.answerMountConfirmation(choice == 0)
|
||||
return m, cmd
|
||||
}
|
||||
m.closeModal()
|
||||
if choice == 1 {
|
||||
|
||||
@@ -271,6 +271,23 @@ func (m Model) viewInner() string {
|
||||
return m.toastOverlay(main)
|
||||
}
|
||||
|
||||
// mountPromptBounds is where the working-directory prompt is drawn. It is placed
|
||||
// by cornerOverlay rather than centered, so a click has to be tested against
|
||||
// these bounds and not the ones the other modals use.
|
||||
func (m Model) mountPromptBounds() (left, top int, panel string) {
|
||||
panel = m.modalView()
|
||||
if panel == "" {
|
||||
return 0, 0, ""
|
||||
}
|
||||
_, _, chatWidth, _ := m.layout()
|
||||
left = max(0, min(chatWidth, m.width)-lipgloss.Width(panel))
|
||||
statusH := 0
|
||||
if m.statusVisible() {
|
||||
statusH = 1
|
||||
}
|
||||
return left, max(0, m.inputTop()-statusH-lipgloss.Height(panel)), panel
|
||||
}
|
||||
|
||||
// cornerOverlay splices a panel in directly above the composer, right-aligned
|
||||
// with it, leaving the rest of the view visible behind it.
|
||||
func (m Model) cornerOverlay(view, panel string) string {
|
||||
|
||||
@@ -220,6 +220,13 @@ func (m Model) confirmView(title string, width int, border, titleColor lipgloss.
|
||||
return m.confirmDialog(title, "", width, border, titleColor, red, "Yes", "No")
|
||||
}
|
||||
|
||||
// The mount prompt's buttons, named so the renderer and the click test cannot
|
||||
// drift apart.
|
||||
const (
|
||||
mountConfirmLabel = "Mount"
|
||||
mountCancelLabel = "Skip"
|
||||
)
|
||||
|
||||
// mountConfirmView asks before a target-less scan mounts the working directory.
|
||||
// It is a compact prompt docked in the corner of the live view: nothing is
|
||||
// prepared until it is answered, and the directory is a workspace rather than a
|
||||
@@ -232,8 +239,8 @@ func (m Model) mountConfirmView() string {
|
||||
}
|
||||
title := render.Bold(amber).Render("△ Mount working directory?")
|
||||
body := render.Col(white).Render(truncatePath(dir, width-4)) + "\n" +
|
||||
render.Dim().Render("writable in the sandbox")
|
||||
return m.cornerPrompt(title, body, width, "Confirm", "Cancel")
|
||||
render.Dim().Render("writable in the sandbox · skip to run without it")
|
||||
return m.cornerPrompt(title, body, width, mountConfirmLabel, mountCancelLabel)
|
||||
}
|
||||
|
||||
// truncatePath keeps the tail of a path visible, which is the part that
|
||||
|
||||
@@ -15,7 +15,7 @@ def _ctx(ctx: RunContextWrapper) -> dict[str, Any]:
|
||||
|
||||
|
||||
@function_tool
|
||||
async def respond_to_user(ctx: RunContextWrapper, message: str) -> str:
|
||||
async def respond_to_user(ctx: RunContextWrapper, message: str = "") -> str:
|
||||
"""Answer the user and hand control back to them.
|
||||
|
||||
This is the ONLY way to yield to the user. Delivering the message and
|
||||
@@ -45,6 +45,10 @@ async def respond_to_user(ctx: RunContextWrapper, message: str) -> str:
|
||||
have followed the tool calls that led here. Lead with the
|
||||
answer or the decision you need, and if you are blocked, say
|
||||
exactly what you need from them.
|
||||
|
||||
Omit it when you have just said your piece as plain text and
|
||||
only need to wait: that text has already reached them, and
|
||||
repeating it makes them read the same answer twice.
|
||||
"""
|
||||
inner = _ctx(ctx)
|
||||
coordinator = coordinator_from_context(inner)
|
||||
|
||||
@@ -110,12 +110,19 @@ def _get_agent_todos(agent_id: str) -> dict[str, dict[str, Any]]:
|
||||
|
||||
|
||||
def _normalize_priority(priority: str | None, default: str = "normal") -> str:
|
||||
candidate = (priority or default or "normal").lower()
|
||||
candidate = str(priority or default or "normal").strip().lower()
|
||||
if candidate not in VALID_PRIORITIES:
|
||||
raise ValueError(f"Invalid priority. Must be one of: {', '.join(VALID_PRIORITIES)}")
|
||||
return candidate
|
||||
|
||||
|
||||
def _coerce_priority(priority: str | None, default: str = "normal") -> str:
|
||||
try:
|
||||
return _normalize_priority(priority, default)
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
|
||||
def _sorted_todos(agent_id: str) -> list[dict[str, Any]]:
|
||||
todos_list = [
|
||||
{**todo, "todo_id": todo_id} for todo_id, todo in _get_agent_todos(agent_id).items()
|
||||
@@ -285,11 +292,16 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str:
|
||||
- ``description`` (str, optional): extra context or
|
||||
acceptance criteria.
|
||||
- ``priority`` (str, optional): one of ``"low"`` /
|
||||
``"normal"`` / ``"high"`` / ``"critical"``. Defaults to
|
||||
``"normal"``.
|
||||
``"normal"`` / ``"high"`` / ``"critical"``. Anything else,
|
||||
including omitting it, falls back to ``"normal"`` rather
|
||||
than failing.
|
||||
|
||||
Example: ``[{"title": "Probe /admin", "priority": "high"},
|
||||
{"title": "Check JWT alg=none"}]``.
|
||||
|
||||
A title already on the list, or repeated within this call, is
|
||||
skipped rather than duplicated; skipped titles come back under
|
||||
``skipped``.
|
||||
"""
|
||||
agent_id = _agent_id_from(ctx)
|
||||
try:
|
||||
@@ -302,13 +314,21 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str:
|
||||
)
|
||||
|
||||
agent_todos = _get_agent_todos(agent_id)
|
||||
seen = {todo["title"].strip().lower() for todo in agent_todos.values()}
|
||||
created: list[dict[str, Any]] = []
|
||||
skipped: list[dict[str, str]] = []
|
||||
for task in tasks:
|
||||
task_priority = _normalize_priority(task.get("priority"))
|
||||
title = task["title"]
|
||||
key = title.lower()
|
||||
if key in seen:
|
||||
skipped.append({"title": title, "reason": "duplicate title"})
|
||||
continue
|
||||
seen.add(key)
|
||||
task_priority = _coerce_priority(task.get("priority"))
|
||||
todo_id = str(uuid.uuid4())[:6]
|
||||
timestamp = datetime.now(UTC).isoformat()
|
||||
agent_todos[todo_id] = {
|
||||
"title": task["title"],
|
||||
"title": title,
|
||||
"description": task.get("description"),
|
||||
"priority": task_priority,
|
||||
"status": "pending",
|
||||
@@ -316,7 +336,7 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str:
|
||||
"updated_at": timestamp,
|
||||
"completed_at": None,
|
||||
}
|
||||
created.append({"todo_id": todo_id, "title": task["title"], "priority": task_priority})
|
||||
created.append({"todo_id": todo_id, "title": title, "priority": task_priority})
|
||||
except (ValueError, TypeError) as e:
|
||||
return json.dumps(
|
||||
{"success": False, "error": f"Failed to create todo: {e}"},
|
||||
@@ -330,6 +350,7 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str:
|
||||
"success": True,
|
||||
"created": created,
|
||||
"created_count": len(created),
|
||||
"skipped": skipped,
|
||||
"todos": _sorted_todos(agent_id),
|
||||
"total_count": len(_get_agent_todos(agent_id)),
|
||||
},
|
||||
|
||||
@@ -70,7 +70,6 @@ async def test_encoded_list_is_decoded_for_an_array_parameter(schema: dict[str,
|
||||
"auth",
|
||||
"Endpoint /admin leaks user data, and session tokens never expire",
|
||||
'"auth"',
|
||||
"",
|
||||
],
|
||||
)
|
||||
async def test_free_form_strings_are_never_split_into_an_array(value: str) -> None:
|
||||
@@ -79,6 +78,29 @@ async def test_free_form_strings_are_never_split_into_an_array(value: str) -> No
|
||||
assert parsed["tags"] == value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("schema", [_ARRAY, _NULLABLE_ARRAY])
|
||||
@pytest.mark.parametrize("value", ["", " "])
|
||||
async def test_empty_string_becomes_an_empty_array(schema: dict[str, Any], value: str) -> None:
|
||||
parsed = await _roundtrip(schema, {"tags": value})
|
||||
|
||||
assert parsed["tags"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_string_becomes_an_empty_object() -> None:
|
||||
parsed = await _roundtrip(_OBJECT, {"modifications": ""})
|
||||
|
||||
assert parsed["modifications"] == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_string_for_a_string_parameter_is_untouched() -> None:
|
||||
parsed = await _roundtrip(_STRING, {"todos": ""})
|
||||
|
||||
assert parsed["todos"] == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_encoded_mapping_is_decoded_for_an_object_parameter() -> None:
|
||||
parsed = await _roundtrip(_OBJECT, {"modifications": '{"method": "POST"}'})
|
||||
|
||||
@@ -1228,3 +1228,40 @@ async def test_wait_kind_survives_a_snapshot_round_trip() -> None:
|
||||
assert restored.wait_kinds["root"] == "user"
|
||||
assert restored.idle_resume_counts["root"] == 1
|
||||
assert await execution._plain_waiting_timeout(restored, "root") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interactive_nudge_offers_waiting_without_repeating() -> None:
|
||||
"""The nudge is the instruction an agent reads when it is stranded here.
|
||||
|
||||
It is where the option to wait on what was already said has to be, not only
|
||||
in the system prompt: an agent that ended a turn on plain text reasons off
|
||||
this text, and without the clause it restates its answer to reach a tool
|
||||
call, so the user reads it twice.
|
||||
|
||||
The clause holds whatever the turn did, because the agent is the one who
|
||||
knows whether it spoke — this fires for a turn that produced no text at all.
|
||||
"""
|
||||
items = await execution._append_tool_required_message(
|
||||
session=None,
|
||||
context={"parent_id": None},
|
||||
attempt=1,
|
||||
limit=3,
|
||||
interactive=True,
|
||||
)
|
||||
|
||||
assert "with no message if you have already said it" in items[0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autonomous_nudge_does_not_offer_the_user() -> None:
|
||||
"""There is nobody attached to an autonomous run to wait for."""
|
||||
items = await execution._append_tool_required_message(
|
||||
session=None,
|
||||
context={"parent_id": None},
|
||||
attempt=1,
|
||||
limit=3,
|
||||
interactive=False,
|
||||
)
|
||||
|
||||
assert "respond_to_user" not in items[0]["content"]
|
||||
|
||||
@@ -64,3 +64,31 @@ async def test_a_message_that_already_arrived_is_taken_instead_of_parking() -> N
|
||||
assert result["wait_outcome"] == "message_arrived"
|
||||
assert result["pending_messages"] == 1
|
||||
assert coordinator.statuses["root"] == "running"
|
||||
|
||||
|
||||
async def _call_without_message(context: dict[str, Any]) -> dict[str, Any]:
|
||||
ctx = ToolContext(
|
||||
context=context,
|
||||
tool_name="respond_to_user",
|
||||
tool_call_id="call-1",
|
||||
tool_arguments="{}",
|
||||
)
|
||||
raw = await respond_to_user.on_invoke_tool(ctx, "{}")
|
||||
return json.loads(raw) # type: ignore[no-any-return]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parks_without_a_message() -> None:
|
||||
"""An agent that has already said its piece as plain text can just wait.
|
||||
|
||||
The nudge is what leaves it here, and while a message was required the only
|
||||
way to stop was to send the same answer a second time.
|
||||
"""
|
||||
context = await _context(interactive=True)
|
||||
|
||||
result = await _call_without_message(context)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["wait_outcome"] == "waiting"
|
||||
assert result["message"] == ""
|
||||
assert context["coordinator"].statuses["root"] == "waiting"
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agents import ModelSettings
|
||||
|
||||
import strix.tools.notes.tools as notes_tools
|
||||
import strix.tools.todo.tools as todo_tools
|
||||
from strix.core import runner
|
||||
from strix.core.agents import AgentCoordinator
|
||||
from strix.runtime import session_manager
|
||||
|
||||
|
||||
def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None:
|
||||
monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path)
|
||||
monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path)
|
||||
monkeypatch.setattr(runner, "setup_scan_logging", lambda _run_dir: lambda: None)
|
||||
monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None)
|
||||
|
||||
settings = types.SimpleNamespace(
|
||||
llm=types.SimpleNamespace(
|
||||
model="openai/gpt-4o",
|
||||
reasoning_effort="high",
|
||||
force_required_tool_choice=False,
|
||||
timeout=300,
|
||||
prompt_cache=True,
|
||||
extra_headers=None,
|
||||
),
|
||||
runtime=types.SimpleNamespace(max_context_images=3),
|
||||
)
|
||||
monkeypatch.setattr(runner, "load_settings", lambda: settings)
|
||||
monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None)
|
||||
monkeypatch.setattr(
|
||||
runner, "uses_chat_completions_tool_schema", lambda _model, _settings: False
|
||||
)
|
||||
monkeypatch.setattr(todo_tools, "hydrate_todos_from_disk", lambda _state_dir: None)
|
||||
monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _state_dir: None)
|
||||
|
||||
async def _create_or_reuse(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
|
||||
return {"client": object(), "session": object(), "caido_client": None}
|
||||
|
||||
async def _cleanup(*_args: Any, **_kwargs: Any) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse)
|
||||
monkeypatch.setattr(session_manager, "cleanup", _cleanup)
|
||||
monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task")
|
||||
monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: "")
|
||||
monkeypatch.setattr(runner, "make_model_settings", lambda *_a, **_k: ModelSettings())
|
||||
monkeypatch.setattr(runner, "build_strix_agent", lambda **_kwargs: object())
|
||||
monkeypatch.setattr(runner, "make_child_factory", lambda **_kwargs: lambda **_k: object())
|
||||
monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object())
|
||||
|
||||
|
||||
def _root_status(coordinator: AgentCoordinator) -> str:
|
||||
roots = [aid for aid, parent in coordinator.parent_of.items() if parent is None]
|
||||
assert len(roots) == 1
|
||||
return coordinator.statuses[roots[0]]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("interrupt", [KeyboardInterrupt, asyncio.CancelledError])
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_interrupt_leaves_the_root_running_for_resume(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Any, interrupt: type[BaseException]
|
||||
) -> None:
|
||||
_wire_runner(monkeypatch, tmp_path)
|
||||
|
||||
async def _interrupt(*_args: Any, **_kwargs: Any) -> None:
|
||||
raise interrupt()
|
||||
|
||||
monkeypatch.setattr(runner, "run_agent_loop", _interrupt)
|
||||
coordinator = AgentCoordinator()
|
||||
|
||||
with pytest.raises(interrupt):
|
||||
await runner.run_strix_scan(
|
||||
scan_config={"targets": [], "scan_mode": "deep"},
|
||||
scan_id="scan-test",
|
||||
image="img",
|
||||
coordinator=coordinator,
|
||||
)
|
||||
|
||||
assert _root_status(coordinator) == "running"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_real_crash_still_marks_root_failed(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Any
|
||||
) -> None:
|
||||
_wire_runner(monkeypatch, tmp_path)
|
||||
|
||||
async def _boom(*_args: Any, **_kwargs: Any) -> None:
|
||||
raise RuntimeError("boom")
|
||||
|
||||
monkeypatch.setattr(runner, "run_agent_loop", _boom)
|
||||
coordinator = AgentCoordinator()
|
||||
|
||||
with pytest.raises(RuntimeError, match="boom"):
|
||||
await runner.run_strix_scan(
|
||||
scan_config={"targets": [], "scan_mode": "deep"},
|
||||
scan_id="scan-test",
|
||||
image="img",
|
||||
coordinator=coordinator,
|
||||
)
|
||||
|
||||
assert _root_status(coordinator) == "failed"
|
||||
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agents import ModelSettings
|
||||
|
||||
import strix.tools.notes.tools as notes_tools
|
||||
import strix.tools.todo.tools as todo_tools
|
||||
from strix.core import runner
|
||||
from strix.core.agents import AgentCoordinator
|
||||
from strix.runtime import session_manager
|
||||
|
||||
|
||||
def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None:
|
||||
monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path)
|
||||
monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path)
|
||||
monkeypatch.setattr(runner, "setup_scan_logging", lambda _run_dir: lambda: None)
|
||||
monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None)
|
||||
|
||||
settings = _settings()
|
||||
monkeypatch.setattr(runner, "load_settings", lambda: settings)
|
||||
monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _s: None)
|
||||
monkeypatch.setattr(runner, "uses_chat_completions_tool_schema", lambda _m, _s: False)
|
||||
monkeypatch.setattr(todo_tools, "hydrate_todos_from_disk", lambda _d: None)
|
||||
monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _d: None)
|
||||
|
||||
async def _create_or_reuse(*_a: Any, **_k: Any) -> dict[str, Any]:
|
||||
return {"client": object(), "session": object(), "caido_client": None}
|
||||
|
||||
async def _cleanup(*_a: Any, **_k: Any) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse)
|
||||
monkeypatch.setattr(session_manager, "cleanup", _cleanup)
|
||||
monkeypatch.setattr(runner, "build_root_task", lambda _c: "task")
|
||||
monkeypatch.setattr(runner, "build_scope_context", lambda _c: "")
|
||||
monkeypatch.setattr(runner, "make_model_settings", lambda *_a, **_k: ModelSettings())
|
||||
monkeypatch.setattr(runner, "build_strix_agent", lambda **_k: object())
|
||||
monkeypatch.setattr(runner, "make_child_factory", lambda **_k: lambda **_kk: object())
|
||||
monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object())
|
||||
|
||||
|
||||
def _settings() -> Any:
|
||||
return types.SimpleNamespace(
|
||||
llm=types.SimpleNamespace(
|
||||
model="openai/gpt-4o",
|
||||
reasoning_effort="high",
|
||||
force_required_tool_choice=False,
|
||||
timeout=300,
|
||||
prompt_cache=True,
|
||||
extra_headers=None,
|
||||
),
|
||||
runtime=types.SimpleNamespace(max_context_images=3),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_live_child_is_settled_before_sessions_close(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Any
|
||||
) -> None:
|
||||
_wire_runner(monkeypatch, tmp_path)
|
||||
coordinator = AgentCoordinator()
|
||||
child_started = asyncio.Event()
|
||||
child_task: dict[str, asyncio.Task[None]] = {}
|
||||
|
||||
async def _root_finishes(**kwargs: Any) -> None:
|
||||
root_id = kwargs["agent_id"]
|
||||
|
||||
async def _child_mid_turn() -> None:
|
||||
child_started.set()
|
||||
await asyncio.sleep(3600)
|
||||
|
||||
await coordinator.register("child", "Child", parent_id=root_id)
|
||||
task = asyncio.create_task(_child_mid_turn())
|
||||
child_task["t"] = task
|
||||
await coordinator.attach_runtime("child", task=task)
|
||||
await child_started.wait()
|
||||
|
||||
monkeypatch.setattr(runner, "run_agent_loop", _root_finishes)
|
||||
|
||||
await runner.run_strix_scan(
|
||||
scan_config={"targets": [], "scan_mode": "deep"},
|
||||
scan_id="scan-test",
|
||||
image="img",
|
||||
coordinator=coordinator,
|
||||
)
|
||||
|
||||
task = child_task["t"]
|
||||
assert task.done(), "the child task was left running past scan teardown"
|
||||
assert task.cancelled(), "the child was not cancelled cleanly on a finish"
|
||||
@@ -0,0 +1,117 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.core.sessions import open_agent_session
|
||||
|
||||
|
||||
def _count_open_fds() -> int | None:
|
||||
for path in (Path("/proc/self/fd"), Path("/dev/fd")):
|
||||
if path.is_dir():
|
||||
return len(list(path.iterdir()))
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sessions_hold_no_descriptors_while_parked(tmp_path: Path) -> None:
|
||||
"""Descriptor use must track live operations, not the number of sessions.
|
||||
|
||||
The SDK keeps a connection per (session, pool thread) open for the session's
|
||||
whole life. An agent parks rather than exits, so its session lives for the
|
||||
scan, and fan-out multiplies those handles until the process runs out of file
|
||||
descriptors (#1018). A session that is not mid-operation should hold none.
|
||||
"""
|
||||
baseline = _count_open_fds()
|
||||
if baseline is None:
|
||||
pytest.skip("no /proc/self/fd or /dev/fd on this platform")
|
||||
|
||||
sessions = [open_agent_session(f"a{i}", tmp_path / f"s{i}.db") for i in range(60)]
|
||||
try:
|
||||
for _ in range(4):
|
||||
await asyncio.gather(
|
||||
*(s.add_items([{"role": "user", "content": "x"}]) for s in sessions)
|
||||
)
|
||||
await asyncio.gather(*(s.get_items() for s in sessions))
|
||||
parked = _count_open_fds()
|
||||
assert parked is not None
|
||||
# 60 parked sessions, yet descriptors are back at the baseline.
|
||||
assert parked - baseline <= 5, f"parked fds grew by {parked - baseline}"
|
||||
finally:
|
||||
for s in sessions:
|
||||
s.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_in_flight_descriptors_track_concurrency_not_session_count(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
baseline = _count_open_fds()
|
||||
if baseline is None:
|
||||
pytest.skip("no /proc/self/fd or /dev/fd on this platform")
|
||||
|
||||
sessions = [open_agent_session(f"a{i}", tmp_path / f"s{i}.db") for i in range(200)]
|
||||
peak = baseline
|
||||
try:
|
||||
|
||||
async def sample() -> None:
|
||||
nonlocal peak
|
||||
for _ in range(500):
|
||||
current = _count_open_fds()
|
||||
if current is not None:
|
||||
peak = max(peak, current)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
async def load() -> None:
|
||||
for _ in range(4):
|
||||
await asyncio.gather(
|
||||
*(s.add_items([{"role": "user", "content": "x"}]) for s in sessions)
|
||||
)
|
||||
|
||||
await asyncio.gather(load(), sample())
|
||||
# 200 sessions, but peak is bounded by the thread pool, well under 200.
|
||||
assert peak - baseline < 100, f"in-flight fds peaked at +{peak - baseline}"
|
||||
finally:
|
||||
for s in sessions:
|
||||
s.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_history_survives_the_per_operation_connection(tmp_path: Path) -> None:
|
||||
session = open_agent_session("agent-1", tmp_path / "agents.db")
|
||||
try:
|
||||
for i in range(30):
|
||||
await session.add_items([{"role": "user", "content": f"m{i}"}])
|
||||
items = [cast("dict[str, Any]", i) for i in await session.get_items()]
|
||||
assert [i["content"] for i in items] == [f"m{i}" for i in range(30)]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_sessions_sharing_one_file_stay_consistent(tmp_path: Path) -> None:
|
||||
db = tmp_path / "shared.db"
|
||||
sessions = [open_agent_session(f"a{i}", db) for i in range(10)]
|
||||
try:
|
||||
await asyncio.gather(
|
||||
*(s.add_items([{"role": "user", "content": s.session_id}]) for s in sessions)
|
||||
)
|
||||
# Each session sees only its own row despite sharing the file.
|
||||
for s in sessions:
|
||||
items = [cast("dict[str, Any]", i) for i in await s.get_items()]
|
||||
assert [i["content"] for i in items] == [s.session_id]
|
||||
finally:
|
||||
for s in sessions:
|
||||
s.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_closed_session_refuses_operations(tmp_path: Path) -> None:
|
||||
session = open_agent_session("agent-1", tmp_path / "agents.db")
|
||||
await session.add_items([{"role": "user", "content": "x"}])
|
||||
session.close()
|
||||
with pytest.raises(RuntimeError, match="closed"):
|
||||
await session.add_items([{"role": "user", "content": "y"}])
|
||||
@@ -0,0 +1,105 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agents.tool_context import ToolContext
|
||||
|
||||
from strix.tools.todo import tools
|
||||
from strix.tools.todo.tools import _coerce_priority, create_todo
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_store() -> Any:
|
||||
tools._todos_storage.clear()
|
||||
yield
|
||||
tools._todos_storage.clear()
|
||||
|
||||
|
||||
async def _create(todos: list[Any], agent_id: str = "root") -> dict[str, Any]:
|
||||
ctx = ToolContext(
|
||||
context={"agent_id": agent_id},
|
||||
tool_name="create_todo",
|
||||
tool_call_id="call-1",
|
||||
tool_arguments="{}",
|
||||
)
|
||||
raw = await create_todo.on_invoke_tool(ctx, json.dumps({"todos": json.dumps(todos)}))
|
||||
return json.loads(raw) # type: ignore[no-any-return]
|
||||
|
||||
|
||||
def test_unknown_priority_falls_back_to_normal() -> None:
|
||||
assert _coerce_priority("medium") == "normal"
|
||||
assert _coerce_priority("urgent") == "normal"
|
||||
assert _coerce_priority("high") == "high"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_bad_priority_no_longer_discards_the_batch() -> None:
|
||||
result = await _create(
|
||||
[
|
||||
{"title": "Recon", "priority": "medium"},
|
||||
{"title": "Probe /admin", "priority": "sky-high"},
|
||||
{"title": "Report"},
|
||||
]
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["created_count"] == 3
|
||||
by_title = {c["title"]: c["priority"] for c in result["created"]}
|
||||
assert by_title["Recon"] == "normal"
|
||||
assert by_title["Probe /admin"] == "normal"
|
||||
assert by_title["Report"] == "normal"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_titles_within_a_batch_are_skipped() -> None:
|
||||
result = await _create(
|
||||
[
|
||||
{"title": "Subdomain enumeration"},
|
||||
{"title": "Content discovery"},
|
||||
{"title": "Subdomain enumeration"},
|
||||
{"title": "content discovery"},
|
||||
]
|
||||
)
|
||||
|
||||
assert result["created_count"] == 2
|
||||
assert {c["title"] for c in result["created"]} == {
|
||||
"Subdomain enumeration",
|
||||
"Content discovery",
|
||||
}
|
||||
assert len(result["skipped"]) == 2
|
||||
assert all(s["reason"] == "duplicate title" for s in result["skipped"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_title_already_on_the_list_is_not_created_again() -> None:
|
||||
await _create([{"title": "Crawl with katana"}])
|
||||
result = await _create([{"title": "crawl with katana"}, {"title": "JS analysis"}])
|
||||
|
||||
assert [c["title"] for c in result["created"]] == ["JS analysis"]
|
||||
assert [s["title"] for s in result["skipped"]] == ["crawl with katana"]
|
||||
assert result["total_count"] == 2
|
||||
|
||||
|
||||
def test_coerce_never_raises() -> None:
|
||||
assert _coerce_priority("nonsense") == "normal"
|
||||
assert _coerce_priority(None) == "normal"
|
||||
assert _coerce_priority("high") == "high"
|
||||
for value in (2, ["high"], {"p": 1}, True):
|
||||
assert _coerce_priority(value) == "normal" # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_string_priority_does_not_fail_the_batch() -> None:
|
||||
result = await _create(
|
||||
[
|
||||
{"title": "Recon", "priority": 2},
|
||||
{"title": "Probe", "priority": ["high"]},
|
||||
{"title": "Report"},
|
||||
]
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["created_count"] == 3
|
||||
assert {c["priority"] for c in result["created"]} == {"normal"}
|
||||
@@ -234,12 +234,11 @@ async def test_confirming_the_mount_starts_the_scan_without_a_target() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_declining_the_mount_returns_to_the_start_screen() -> None:
|
||||
started = False
|
||||
async def test_declining_the_mount_runs_without_one() -> None:
|
||||
started: list[bool] = []
|
||||
|
||||
async def start(_verify: bool = True) -> None:
|
||||
nonlocal started
|
||||
started = True
|
||||
async def start(verify: bool = True) -> None:
|
||||
started.append(verify)
|
||||
|
||||
os.environ["STRIX_LLM"] = "anthropic/claude-sonnet-4"
|
||||
os.environ["ANTHROPIC_API_KEY"] = "test-key"
|
||||
@@ -250,14 +249,34 @@ async def test_declining_the_mount_returns_to_the_start_screen() -> None:
|
||||
result = await controller.handle("setup.confirm_mount", {"approved": False})
|
||||
|
||||
assert result == {"approved": False}
|
||||
# Nothing was prepared, so the session goes back to the start screen and can
|
||||
# be launched again.
|
||||
assert started is False
|
||||
# Declining skips the directory; it does not abandon the scan.
|
||||
assert started == [False]
|
||||
assert controller.workspace_mount is None
|
||||
assert controller.pending_workspace_mount is None
|
||||
assert controller.setup_mode is True
|
||||
assert controller.scan_started is False
|
||||
assert controller.scan_state == "setup"
|
||||
assert controller.setup_mode is False
|
||||
assert controller.scan_started is True
|
||||
assert controller.scan_state == "running"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_approving_the_mount_runs_with_it() -> None:
|
||||
started: list[bool] = []
|
||||
|
||||
async def start(verify: bool = True) -> None:
|
||||
started.append(verify)
|
||||
|
||||
os.environ["STRIX_LLM"] = "anthropic/claude-sonnet-4"
|
||||
os.environ["ANTHROPIC_API_KEY"] = "test-key"
|
||||
loader._cached = None
|
||||
controller = TuiController(args(), on_start=start)
|
||||
await controller.handle("setup.start", {"verify": False, "mount_working_dir": True})
|
||||
|
||||
result = await controller.handle("setup.confirm_mount", {"approved": True})
|
||||
|
||||
assert result == {"approved": True}
|
||||
assert started == [False]
|
||||
assert controller.workspace_mount == str(Path.cwd())
|
||||
assert controller.scan_state == "running"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -7,19 +7,26 @@ shows what the user actually typed; resuming has to match that.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import json
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.core import execution
|
||||
from strix.core.paths import runtime_state_dir
|
||||
from strix.interface.tui.backend.live_view import TuiLiveView as GoTuiLiveView
|
||||
from strix.interface.tui.live_view import TuiLiveView, _is_internal_agent_turn
|
||||
from strix.interface.tui.live_view import (
|
||||
_INTERNAL_TURN_PREFIXES,
|
||||
TuiLiveView,
|
||||
_is_internal_agent_turn,
|
||||
)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
|
||||
|
||||
def _write_run(run_dir: Path, items: list[dict[str, Any]], agent_id: str = "root") -> None:
|
||||
@@ -176,11 +183,60 @@ def test_internal_turn_classifier_matches_every_injected_form() -> None:
|
||||
"[CRITICAL] Turn budget: 480/500 used (96%).",
|
||||
"== Inherited context from parent (background only) ==",
|
||||
"Your previous message ended a turn without a tool call.",
|
||||
"Your previous response ended the autonomous Strix run without a lifecycle tool call.",
|
||||
"Your previous response ended the autonomous run without a lifecycle tool call.",
|
||||
):
|
||||
assert _is_internal_agent_turn(content), content
|
||||
|
||||
|
||||
def _injected_strings(module: ModuleType) -> list[str]:
|
||||
"""Every string a module can inject, and nothing it merely mentions.
|
||||
|
||||
Parsing rather than searching the text keeps comments out of it, so a stale
|
||||
copy of a message left in a comment cannot pass for the message itself. It
|
||||
also joins adjacent literals for free, which the line wrapping needs, and
|
||||
docstrings are dropped because they describe the code rather than run in it.
|
||||
"""
|
||||
tree = ast.parse(Path(module.__file__ or "").read_text(encoding="utf-8"))
|
||||
docstrings = set()
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Module | ast.ClassDef | ast.FunctionDef | ast.AsyncFunctionDef):
|
||||
continue
|
||||
first = node.body[0] if node.body else None
|
||||
if isinstance(first, ast.Expr) and isinstance(first.value, ast.Constant):
|
||||
docstrings.add(id(first.value))
|
||||
|
||||
literals: list[str] = []
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Constant):
|
||||
if isinstance(node.value, str) and id(node) not in docstrings:
|
||||
literals.append(node.value)
|
||||
elif isinstance(node, ast.JoinedStr):
|
||||
literals.append(
|
||||
"".join(
|
||||
part.value
|
||||
for part in node.values
|
||||
if isinstance(part, ast.Constant) and isinstance(part.value, str)
|
||||
)
|
||||
)
|
||||
return literals
|
||||
|
||||
|
||||
def test_internal_turn_prefixes_still_match_what_is_injected() -> None:
|
||||
"""The classifier copies sentences out of another module, so they can drift.
|
||||
|
||||
Both nudges are written inline in strix.core.execution, so there is nothing to
|
||||
import and compare against. Read them back out of what that module can inject.
|
||||
"""
|
||||
injected = _injected_strings(execution)
|
||||
nudges = [prefix for prefix in _INTERNAL_TURN_PREFIXES if prefix.startswith("Your previous")]
|
||||
assert nudges, "the no-tool-call nudges are no longer in the classifier"
|
||||
for nudge in nudges:
|
||||
assert any(nudge in literal for literal in injected), (
|
||||
f"the classifier expects {nudge!r}, which strix.core.execution no longer "
|
||||
f"injects. A resumed scan would show that nudge as the user's own message."
|
||||
)
|
||||
|
||||
|
||||
def test_internal_turn_classifier_keeps_bracketed_user_text() -> None:
|
||||
"""A leading bracket is not enough: typed text often starts with one."""
|
||||
for content in (
|
||||
|
||||
Reference in New Issue
Block a user