refactor(agent/display,error_classifier): simplify shell tokenizers, share abort-fallback verdict hints

This commit is contained in:
Teknium
2026-09-02 19:19:34 -07:00
parent 5f37b25c88
commit b4b3baeda1
2 changed files with 34 additions and 53 deletions

View File

@@ -210,18 +210,13 @@ def _scan_quoted(text: str) -> Iterator[tuple[int, str, bool]]:
def _split_shell_words(segment: str) -> list[str]:
words: list[str] = []
buf: list[str] = []
parts: list[list[str]] = [[]]
for _, ch, quoted in _scan_quoted(segment):
if not quoted and ch.isspace():
if buf:
words.append("".join(buf))
buf = []
parts.append([])
else:
buf.append(ch)
if buf:
words.append("".join(buf))
return words
parts[-1].append(ch)
return ["".join(p) for p in parts if p]
def _strip_shell_pipe_tail(segment: str) -> str:
@@ -236,30 +231,20 @@ def _strip_shell_pipe_tail(segment: str) -> str:
def _split_shell_compound(command: str) -> list[str]:
"""Split on unquoted ``&&`` / ``||`` / ``;`` / newline, dropping pipe tails per segment."""
segments: list[str] = []
buf: list[str] = []
raw: list[list[str]] = [[]]
skip = False
def _flush() -> None:
segment = _strip_shell_pipe_tail("".join(buf).strip())
if segment:
segments.append(segment)
buf.clear()
for i, ch, quoted in _scan_quoted(command):
if skip:
skip = False
continue
if not quoted and (command.startswith("&&", i) or command.startswith("||", i)):
_flush()
elif not quoted and (command.startswith("&&", i) or command.startswith("||", i)):
raw.append([])
skip = True
continue
if not quoted and ch in {";", "\n"}:
_flush()
continue
buf.append(ch)
_flush()
return segments
elif not quoted and ch in {";", "\n"}:
raw.append([])
else:
raw[-1].append(ch)
segments = (_strip_shell_pipe_tail("".join(buf).strip()) for buf in raw)
return [s for s in segments if s]
def _shell_head_word(segment: str) -> str:
@@ -278,13 +263,12 @@ def _clean_shell_segment(segment: str) -> str:
while i < len(words):
word = words[i]
if re.match(r"^\d*(?:>>?|<)$", word):
i += 2
continue
if re.match(r"^\d*(?:>&|<&)\d+$", word) or re.match(r"^\d*>&\d+$", word):
i += 2 # operator + target
elif re.match(r"^\d*(?:>&|<&)\d+$", word):
i += 1
else:
out.append(word)
i += 1
continue
out.append(word)
i += 1
return " ".join(out).strip()

View File

@@ -323,26 +323,23 @@ def _v(reason: FailoverReason, **hints: Any) -> Verdict:
_ROTATE_FALLBACK = {"should_rotate_credential": True, "should_fallback": True}
_ABORT_FALLBACK = {"retryable": False, "should_fallback": True}
_R = FailoverReason
_V_BILLING = _v(FailoverReason.billing, retryable=False, **_ROTATE_FALLBACK)
_V_RATE_LIMIT = _v(FailoverReason.rate_limit, **_ROTATE_FALLBACK)
_V_OVERLOADED = _v(FailoverReason.overloaded)
_V_SERVER_ERROR = _v(FailoverReason.server_error)
_V_CONTEXT_OVERFLOW = _v(FailoverReason.context_overflow, should_compress=True)
_V_PAYLOAD_TOO_LARGE = _v(FailoverReason.payload_too_large, should_compress=True)
_V_MODEL_NOT_FOUND = _v(FailoverReason.model_not_found, retryable=False, should_fallback=True)
_V_POLICY_BLOCKED = _v(FailoverReason.provider_policy_blocked, retryable=False)
_V_CONTENT_BLOCKED = _v(FailoverReason.content_policy_blocked, retryable=False, should_fallback=True)
_V_FORMAT_ERROR = _v(FailoverReason.format_error, retryable=False, should_fallback=True)
_V_AUTH_ROTATE = _v(FailoverReason.auth, retryable=False, **_ROTATE_FALLBACK)
_V_AUTH_FALLBACK = _v(FailoverReason.auth, retryable=False, should_fallback=True)
_V_TIMEOUT = _v(FailoverReason.timeout)
_V_SSL_CERT = _v(FailoverReason.ssl_cert_verification, retryable=False)
_V_IMAGE_TOO_LARGE = _v(FailoverReason.image_too_large)
_V_IMAGE_CORRUPT = _v(FailoverReason.image_corrupt)
_V_MULTIMODAL = _v(FailoverReason.multimodal_tool_content_unsupported)
_V_INVALID_ENCRYPTED = _v(FailoverReason.invalid_encrypted_content)
_V_UNKNOWN = _v(FailoverReason.unknown)
_V_BILLING = _v(_R.billing, retryable=False, **_ROTATE_FALLBACK)
_V_RATE_LIMIT = _v(_R.rate_limit, **_ROTATE_FALLBACK)
_V_AUTH_ROTATE = _v(_R.auth, retryable=False, **_ROTATE_FALLBACK)
_V_AUTH_FALLBACK = _v(_R.auth, **_ABORT_FALLBACK)
_V_MODEL_NOT_FOUND = _v(_R.model_not_found, **_ABORT_FALLBACK)
_V_CONTENT_BLOCKED = _v(_R.content_policy_blocked, **_ABORT_FALLBACK)
_V_FORMAT_ERROR = _v(_R.format_error, **_ABORT_FALLBACK)
_V_POLICY_BLOCKED = _v(_R.provider_policy_blocked, retryable=False)
_V_SSL_CERT = _v(_R.ssl_cert_verification, retryable=False)
_V_CONTEXT_OVERFLOW = _v(_R.context_overflow, should_compress=True)
_V_PAYLOAD_TOO_LARGE = _v(_R.payload_too_large, should_compress=True)
_V_OVERLOADED, _V_SERVER_ERROR, _V_TIMEOUT, _V_UNKNOWN = map(_v, (_R.overloaded, _R.server_error, _R.timeout, _R.unknown))
_V_IMAGE_TOO_LARGE, _V_IMAGE_CORRUPT = _v(_R.image_too_large), _v(_R.image_corrupt)
_V_MULTIMODAL, _V_INVALID_ENCRYPTED = _v(_R.multimodal_tool_content_unsupported), _v(_R.invalid_encrypted_content)
def _billing_hints(error_msg: str) -> Verdict: