feat(llm): detect refusal/deflection responses and retry with a new seed (#36)

Some models (observed with an abliterated Qwen variant) answer a request
they judge "politically sensitive" with hedging refusal language instead of
erroring or returning blank — neither of which the existing blank-response
or schema-validation retry triggers catch, so the refusal just passed
through as a normal result.

Adds a hybrid detector to _llm/retry.py: a free lexical/regex pass catches
blatant refusals ("I cannot generate...") without any network call; a
response that's short and/or hedge-y enough to be ambiguous additionally
gets an embedding-similarity check against canonical refusal exemplars, via
a new optional LLMProvider.embed() (Ollama's /api/embed, llama-server's
/v1/embeddings) — deliberately model-agnostic to the backend, since this is
a model-behavior concern, not a backend one. A detected refusal is treated
exactly like a blank response or validation failure: retried with a bumped
seed (existing next_seed()/RETRY_BACKOFF_SECS machinery), never returned to
the caller.

Wired into all four retry loops: OllamaProvider.chat()/chat_structured(),
LlamaCppProvider.chat(), and the shared _llm/chat.py chat_structured() used
by LlamaCppProvider.chat_structured(). New OllamaOptionRefusalRetry node
(same composable OLLAMA_OPTIONS chain as OllamaOptionDisableThinking)
configures it: enabled toggle, optional embedding_model (blank = lexical-
only), similarity threshold.


Claude-Session: https://claude.ai/code/session_01YArD9ZjBWKsAvazmS48amA

Co-authored-by: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
James Veitch
2026-07-29 17:11:08 +01:00
committed by GitHub
co-authored by Claude Sonnet 5
parent 2724f535f5
commit 1f138e5991
12 changed files with 1150 additions and 11 deletions
+3
View File
@@ -17,6 +17,7 @@ from .ollama import (
OllamaOptionDisableThinking,
OllamaOptionExtraBody,
OllamaOptionMaxTokens,
OllamaOptionRefusalRetry,
OllamaOptionRepeatPenalty,
OllamaOptionSeed,
OllamaOptionTemperature,
@@ -48,6 +49,7 @@ NODE_CLASS_MAPPINGS = {
"OllamaOptionTopK": OllamaOptionTopK,
"OllamaOptionRepeatPenalty": OllamaOptionRepeatPenalty,
"OllamaOptionDisableThinking": OllamaOptionDisableThinking,
"OllamaOptionRefusalRetry": OllamaOptionRefusalRetry,
"OllamaOptionExtraBody": OllamaOptionExtraBody,
"OllamaDebugHistory": OllamaDebugHistory,
"OllamaHistoryLength": OllamaHistoryLength,
@@ -75,6 +77,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"OllamaOptionTopK": "Ollama Option — Top K",
"OllamaOptionRepeatPenalty": "Ollama Option — Repeat Penalty",
"OllamaOptionDisableThinking": "Ollama Option — Disable Thinking",
"OllamaOptionRefusalRetry": "Ollama Option — Refusal Retry",
"OllamaOptionExtraBody": "Ollama Option — Extra Body",
"OllamaDebugHistory": "Ollama Debug History",
"OllamaHistoryLength": "Ollama History Length",
+36 -2
View File
@@ -46,7 +46,7 @@ from pydantic_ai.providers.openai import OpenAIProvider
from pydantic_ai.settings import ModelSettings
from .provider import Message
from .retry import RETRY_BACKOFF_SECS, next_seed
from .retry import RETRY_BACKOFF_SECS, EmbedFn, is_refusal, next_seed
_STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS = (
UnexpectedModelBehavior,
@@ -126,6 +126,7 @@ async def chat_structured(
options: dict | None = None,
max_retries: int = 2,
timeout_secs: float = 300.0,
embed_fn: EmbedFn | None = None,
) -> BaseModel:
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
return a validated instance of ``schema``.
@@ -155,6 +156,15 @@ async def chat_structured(
Retries up to ``max_retries`` times (clamped 0-5) on validation failure
before raising ``RuntimeError``. Never returns a value that failed
validation against ``schema``.
``options`` may also carry a ``"refusal_retry"`` config dict (same
comfydv-level convention as ``"think"``, emitted by
``OllamaOptionRefusalRetry``) — a detected refusal/deflection (see
``_llm/retry.py``) is treated exactly like a validation failure: retried
with a bumped seed rather than returned to the caller. ``embed_fn`` is
``LlamaCppProvider``'s own ``embed()``, bound to whatever embedding
model the config names — passed in rather than looked up here since
this module has no provider instance of its own to call.
"""
if not messages or messages[-1].role != "user":
raise ValueError(
@@ -175,6 +185,11 @@ async def chat_structured(
options = dict(options)
think = options.pop("think")
options = options or None
refusal_cfg = None
if options and "refusal_retry" in options:
options = dict(options)
refusal_cfg = options.pop("refusal_retry")
options = options or None
extra_body: dict = {}
if options:
extra_body["options"] = options
@@ -233,7 +248,26 @@ async def chat_structured(
# not a static type parameter), so the checker can't narrow
# result.output past Agent's default `str` — cast to the
# function's declared return type, which schema is a subtype of.
return cast(BaseModel, result.output)
output = cast(BaseModel, result.output)
if refusal_cfg and refusal_cfg.get("enabled"):
# Re-serialized, not the original wire text — pydantic-ai's
# NativeOutput doesn't expose that separately, and the
# regex/embedding check works the same either way (same
# textual content, just re-encoded).
content = output.model_dump_json()
refused = await is_refusal(
content,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
)
if refused:
last_error = RuntimeError("refusal/deflection detected")
last_invalid_text = content
if attempt < total_attempts:
await asyncio.sleep(RETRY_BACKOFF_SECS)
continue
return output
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
last_error = exc
last_invalid_text = str(exc)
+71 -4
View File
@@ -18,9 +18,16 @@ import logging
from pydantic import BaseModel
from .ollama_provider import _TTLLRUCache, _cache_key, _get_json, _pop_think, _post_json
from .ollama_provider import (
_TTLLRUCache,
_cache_key,
_get_json,
_pop_refusal_retry,
_pop_think,
_post_json,
)
from .provider import Message, ModelInfo, ModelStatus
from .retry import RETRY_BACKOFF_SECS, next_seed
from .retry import RETRY_BACKOFF_SECS, is_refusal, next_seed
logger = logging.getLogger(__name__)
@@ -186,6 +193,15 @@ class LlamaCppProvider:
) -> str:
payload_messages = [_to_openai_message(m) for m in messages]
options, think = _pop_think(options)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
if (
refusal_cfg
and refusal_cfg.get("enabled")
and refusal_cfg.get("embedding_model")
):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
total_attempts = max(0, min(int(max_retries), 5)) + 1
response_text = ""
@@ -246,8 +262,17 @@ class LlamaCppProvider:
else ""
)
if response_text.strip():
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
refused = False
if refusal_cfg and refusal_cfg.get("enabled"):
refused = await is_refusal(
response_text,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
)
if not refused:
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
if attempt < total_attempts:
await asyncio.sleep(RETRY_BACKOFF_SECS)
@@ -282,6 +307,16 @@ class LlamaCppProvider:
if hit:
return schema.model_validate(cached)
embed_fn = None
refusal_cfg = (options or {}).get("refusal_retry")
if (
refusal_cfg
and refusal_cfg.get("enabled")
and refusal_cfg.get("embedding_model")
):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
result = await _chat_structured_impl(
base_url=f"{self.host}/v1",
model=model,
@@ -291,6 +326,38 @@ class LlamaCppProvider:
options=options,
max_retries=max_retries,
timeout_secs=timeout_secs,
embed_fn=embed_fn,
)
_CHAT_RESPONSE_CACHE.set(cache_key, result.model_dump())
return result
async def embed(self, model: str, text: str) -> list[float] | None:
"""POST {host}/v1/embeddings — llama-server's OpenAI-compatible
embeddings endpoint.
Requires an embedding-capable model to be loaded in the router
(typically a *different* model from whatever's answering chat
requests) — not live-verified against a running llama-server (no
instance available at implementation time), mirroring this
provider's other sourced-from-docs-not-verified caveats. Returns
``None`` rather than raising on any failure, same contract as
``OllamaProvider.embed()``.
"""
if not model.strip() or not text.strip():
return None
try:
result = await _post_json(
f"{self.host}/v1/embeddings",
{"model": model, "input": text},
timeout=30.0,
headers=self.headers,
)
except Exception:
return None
data = result.get("data")
if not isinstance(data, list) or not data:
return None
vec = data[0].get("embedding")
if not isinstance(vec, list) or not vec:
return None
return vec
+97 -3
View File
@@ -19,7 +19,7 @@ import time
from pydantic import BaseModel, ValidationError
from .provider import Message, ModelInfo, ModelStatus
from .retry import RETRY_BACKOFF_SECS, next_seed
from .retry import RETRY_BACKOFF_SECS, is_refusal, next_seed
logger = logging.getLogger(__name__)
@@ -102,6 +102,22 @@ def _pop_think(options: dict | None) -> tuple[dict | None, bool | None]:
return (remaining or None), think
def _pop_refusal_retry(options: dict | None) -> tuple[dict | None, dict | None]:
"""Split a ``"refusal_retry"`` config dict out of a generic ``options``
dict — same convention as ``_pop_think``: ``OllamaOptionRefusalRetry``
merges ``{"refusal_retry": {"enabled", "embedding_model", "threshold"}}``
into the same composable ``OLLAMA_OPTIONS`` chain every other
``OllamaOption*`` node feeds into ``ChatCompletion``'s ``options``
input, and neither Ollama's nor llama.cpp's own API recognizes this key,
so every provider pops it out here before building its request.
"""
if not options or "refusal_retry" not in options:
return options, None
remaining = dict(options)
cfg = remaining.pop("refusal_retry")
return (remaining or None), cfg
_MODEL_LIST_CACHE = _TTLLRUCache(maxsize=32, ttl_seconds=20.0)
_CHAT_RESPONSE_CACHE = _TTLLRUCache(maxsize=64, ttl_seconds=None)
_CAPABILITY_CACHE = _TTLLRUCache(maxsize=32, ttl_seconds=300.0)
@@ -353,6 +369,15 @@ class OllamaProvider:
# array (ADR-008 — no transform needed for /api/chat).
payload_messages = [m.model_dump(exclude_none=True) for m in messages]
options, think = _pop_think(options)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
if (
refusal_cfg
and refusal_cfg.get("enabled")
and refusal_cfg.get("embedding_model")
):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
total_attempts = max(0, min(int(max_retries), 5)) + 1
response_text = ""
incomplete = False
@@ -395,8 +420,21 @@ class OllamaProvider:
)
response_text = result.get("message", {}).get("content", "")
if response_text.strip():
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
refused = False
if refusal_cfg and refusal_cfg.get("enabled"):
refused = await is_refusal(
response_text,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
)
if not refused:
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
return response_text
# A detected refusal is handled exactly like a blank
# response below: fall through to the backoff/retry with a
# bumped seed (next_seed), rather than returning the refusal
# text to the caller.
# done: false alongside blank content is a distinct signal from
# an ordinary blank generation — it's Ollama answering before
@@ -457,6 +495,15 @@ class OllamaProvider:
payload_messages = [m.model_dump(exclude_none=True) for m in messages]
json_schema = schema.model_json_schema()
options, think = _pop_think(options)
options, refusal_cfg = _pop_refusal_retry(options)
embed_fn = None
if (
refusal_cfg
and refusal_cfg.get("enabled")
and refusal_cfg.get("embedding_model")
):
embedding_model = refusal_cfg["embedding_model"]
embed_fn = lambda t: self.embed(embedding_model, t) # noqa: E731
cache_key = _cache_key(
"chat_structured",
self.host,
@@ -517,6 +564,25 @@ class OllamaProvider:
await asyncio.sleep(RETRY_BACKOFF_SECS)
continue
if refusal_cfg and refusal_cfg.get("enabled"):
# Checked against the raw JSON text, not a specific parsed
# field: ChatCompletion's schema is caller-defined and this
# provider has no idea which field would carry refusal
# language — the regex/embedding check still matches text
# sitting inside a JSON string value either way.
refused = await is_refusal(
content,
embed_fn=embed_fn,
embed_cache_key=refusal_cfg.get("embedding_model", ""),
threshold=refusal_cfg.get("threshold", 0.82),
)
if refused:
last_error = RuntimeError("refusal/deflection detected")
last_invalid_text = content
if attempt < total_attempts:
await asyncio.sleep(RETRY_BACKOFF_SECS)
continue
_CHAT_RESPONSE_CACHE.set(cache_key, parsed.model_dump())
return parsed
@@ -525,3 +591,31 @@ class OllamaProvider:
f"{total_attempts} attempt(s) (model={model!r}). Last error: "
f"{last_error}. Last response (truncated): {last_invalid_text[:300]!r}"
)
async def embed(self, model: str, text: str) -> list[float] | None:
"""POST {host}/api/embed — Ollama's native embeddings endpoint.
Returns ``None`` rather than raising on any failure (wrong/missing
embedding model, unreachable server, malformed response) — this is
a best-effort capability per the ``LLMProvider`` protocol, and its
one current caller (refusal-retry detection) already treats
``None`` as "skip the embedding check", not an error.
"""
if not model.strip() or not text.strip():
return None
try:
result = await _post_json(
f"{self.host}/api/embed",
{"model": model, "input": text},
timeout=30.0,
headers=self.headers,
)
except Exception:
return None
embeddings = result.get("embeddings")
if not isinstance(embeddings, list) or not embeddings:
return None
vec = embeddings[0]
if not isinstance(vec, list) or not vec:
return None
return vec
+15
View File
@@ -124,3 +124,18 @@ class LLMProvider(Protocol):
same per-provider translation.
"""
...
async def embed(self, model: str, text: str) -> list[float] | None:
"""Embedding vector for ``text``, or ``None`` if unavailable.
Best-effort, not a core capability every deployment has configured:
``model`` must itself be embedding-capable, which is typically a
*different* model from whatever's answering chat requests (e.g.
``nomic-embed-text``, not the model passed to ``chat()``). Returns
``None`` rather than raising when embeddings aren't usable right now
(wrong/missing model, unreachable server) — the one current caller,
refusal-retry detection (see ``_llm/retry.py``), degrades gracefully
to lexical-only detection when this returns ``None``, so a provider
with no embedding model configured is never a hard failure.
"""
...
+167
View File
@@ -11,8 +11,21 @@ each re-deriving the policy.
Only blank/whitespace-only responses trigger a retry for plain chat() —
not merely "short" ones — because a fixed length threshold would misfire on
legitimately short, valid answers (single-word replies, labels, "yes"/"no").
Refusal/deflection detection (below) is a separate, opt-in trigger for the
same retry-with-a-new-seed mechanism: some models (observed with an
abliterated Qwen variant) answer with a soft refusal on a topic they judge
"sensitive" instead of erroring or returning blank, so neither of the above
checks catches it. This is deliberately a model-behavior concern, not a
backend one — every ``LLMProvider`` implementation (Ollama, llama.cpp, and
whatever comes next) wires the same detector into its own retry loop via its
own ``embed()``, rather than each backend inventing its own heuristic.
"""
import math
import re
from collections.abc import Awaitable, Callable
RETRY_BACKOFF_SECS = 1.5
"""Flat delay between retries — gives a still-loading model time to finish
before the next attempt, rather than hammering it with identical requests
@@ -32,3 +45,157 @@ def next_seed(options: dict | None, attempt: int) -> int:
if options and isinstance(options.get("seed"), int):
base = options["seed"]
return base + (attempt - 1)
# ---------------------------------------------------------------------------
# Refusal/deflection detection
# ---------------------------------------------------------------------------
#
# Hybrid, cheapest-check-first: a fast, free lexical pass catches the blatant
# majority ("I cannot generate...") without ever touching the network; only
# a response that's short and/or hedge-y enough to be genuinely ambiguous
# pays for an embedding call. A long, on-topic response never reaches the
# embedding step at all.
REFUSAL_LEXICAL_PATTERNS: tuple[re.Pattern, ...] = tuple(
re.compile(p, re.IGNORECASE)
for p in (
r"\bI\s*(?:'m|\s+am)?\s*(?:cannot|can't|won't|will not)\b[^.]{0,60}?\b"
r"(?:generate|create|produce|write|provide|help|assist|describe|depict|continue)\b",
r"\bI'm (?:not able|unable) to\b",
r"\bI don't feel comfortable\b",
r"\bI'm sorry,?\s*(?:but\s+)?I\s*(?:can't|cannot)\b",
r"\bas an AI\b[^.]{0,60}?\b(?:cannot|can't|unable|not able)\b",
r"\bI must (?:decline|refuse)\b",
r"\bthis (?:request|content|image|prompt) (?:violates|goes against)\b",
r"\bnot something I(?:'m| am)? (?:able to|going to) (?:help with|generate|create|provide)\b",
r"\bI (?:won't|will not) be (?:able to )?(?:generat|creat|provid)",
)
)
_SOFT_HEDGE_KEYWORDS: tuple[str, ...] = (
"sensitive",
"unfortunately",
"i must",
"policy",
"guideline",
"responsible ai",
"not appropriate",
"instead, i",
"i'd rather",
"i would rather",
)
_AMBIGUOUS_LENGTH_THRESHOLD = 400
"""Below this many characters, a response is short enough that a soft
refusal is plausible and worth the extra embedding check — chosen well
under this pipeline's normal structured-JSON response sizes (typically
1000+ characters), not a tuned/validated threshold."""
REFUSAL_EXEMPLARS: tuple[str, ...] = (
"I cannot generate an image or description involving this topic.",
"I'm not able to help with this request as it involves a politically sensitive subject.",
"I don't feel comfortable creating content related to this figure or event.",
"This request touches on a sensitive political topic that I must decline to engage with.",
"As an AI, I'm unable to produce content depicting this scenario.",
"I'd rather not generate this — it involves content I'm not comfortable with.",
)
def is_lexical_refusal(text: str) -> bool:
"""Cheap, free regex pass — catches the blatant majority of refusals."""
return any(p.search(text) for p in REFUSAL_LEXICAL_PATTERNS)
def is_ambiguous(text: str) -> bool:
"""Whether ``text`` is short/hedge-y enough to be worth the pricier
embedding check, having already failed the free lexical pass.
Deliberately cheap and approximate — false positives here only cost one
extra embedding call, false negatives skip a refusal that a real
similarity check might have caught. Not meant to be a precise signal on
its own, just a gate on when the more expensive check runs at all.
"""
stripped = text.strip()
if not stripped:
return False
if len(stripped) < _AMBIGUOUS_LENGTH_THRESHOLD:
return True
lowered = stripped.lower()
return any(keyword in lowered for keyword in _SOFT_HEDGE_KEYWORDS)
def cosine_similarity(a: list[float], b: list[float]) -> float:
"""Standard cosine similarity, no numpy dependency (comfydv has none)."""
if not a or not b or len(a) != len(b):
return 0.0
dot = sum(x * y for x, y in zip(a, b))
norm_a = math.sqrt(sum(x * x for x in a))
norm_b = math.sqrt(sum(y * y for y in b))
if norm_a == 0.0 or norm_b == 0.0:
return 0.0
return dot / (norm_a * norm_b)
EmbedFn = Callable[[str], Awaitable[list[float] | None]]
_exemplar_embedding_cache: dict[str, list[list[float]]] = {}
async def _exemplar_embeddings(embed_fn: EmbedFn, cache_key: str) -> list[list[float]]:
"""Embed ``REFUSAL_EXEMPLARS`` once per ``cache_key`` (an embedding-model
identifier) and reuse — the exemplars never change, only the embedding
space (i.e. which model produced the vectors) does.
"""
cached = _exemplar_embedding_cache.get(cache_key)
if cached is not None:
return cached
embeddings = []
for exemplar in REFUSAL_EXEMPLARS:
vec = await embed_fn(exemplar)
if not vec:
# An embedding call failing for one exemplar almost certainly
# means embeddings aren't usable at all right now (wrong/missing
# embedding model, unreachable server) — bail out rather than
# caching a partial, unusable exemplar set.
return []
embeddings.append(vec)
_exemplar_embedding_cache[cache_key] = embeddings
return embeddings
async def is_refusal(
text: str,
*,
embed_fn: EmbedFn | None = None,
embed_cache_key: str = "",
threshold: float = 0.82,
) -> bool:
"""Hybrid refusal/deflection detector: free lexical pass first, then an
optional embedding-similarity fallback only for ambiguous responses.
``embed_fn`` is supplied by the caller's own ``LLMProvider.embed()`` —
this function has no idea which backend or model produced ``text``, by
design (ADR: refusal detection is a model-behavior concern, not a
backend one). ``embed_fn=None`` (no embedding model configured) degrades
to lexical-only detection rather than erroring. Any failure while
embedding (unreachable server, no embedding-capable model loaded) is
swallowed the same way — an optional enhancement failing shouldn't take
down the retry loop it's assisting.
"""
if not text or not text.strip():
return False # blank responses are the *other* retry trigger, not this one
if is_lexical_refusal(text):
return True
if embed_fn is None or not is_ambiguous(text):
return False
try:
exemplar_vecs = await _exemplar_embeddings(embed_fn, embed_cache_key)
if not exemplar_vecs:
return False
text_vec = await embed_fn(text)
if not text_vec:
return False
except Exception:
return False
return max(cosine_similarity(text_vec, vec) for vec in exemplar_vecs) >= threshold
+81
View File
@@ -914,6 +914,87 @@ class OllamaOptionDisableThinking:
return (_merge_option(options, "think", not disable_thinking),)
class OllamaOptionRefusalRetry:
"""Retry with a bumped seed when a response looks like a soft
refusal/deflection rather than an actual error (see
``_llm/retry.py``'s ``is_refusal()``).
Observed with an abliterated Qwen variant: it sometimes answers a
request it judges "politically sensitive" with hedging refusal
language instead of erroring or returning blank — neither of which the
existing blank-response or schema-validation retry triggers catch, so
without this the response just passes through as-is.
Rides the same composable ``OLLAMA_OPTIONS`` chain as every other
``OllamaOption*`` node, but like ``OllamaOptionDisableThinking``, the
``"refusal_retry"`` key this node emits is a comfydv-level convention,
not an Ollama-native sampling param: every ``LLMProvider``
implementation pops it out of ``options`` and drives its own retry
loop with it (ADR: refusal detection is a model-behavior concern, not
a backend one — see ``LLMProvider.embed()`` in ``_llm/provider.py``).
Works for both backends from the same node.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"enabled": ("BOOLEAN", {"default": True}),
"embedding_model": (
"STRING",
{
"default": "",
"tooltip": (
"Name of a separate embedding-capable model "
'(e.g. "nomic-embed-text") — NOT the chat '
"model itself; most chat models can't produce "
"usable embeddings. Leave blank to skip the "
"embedding-similarity check and detect only "
"blatant, literal refusal phrases (still "
"useful, cheaper, catches less)."
),
},
),
"threshold": (
"FLOAT",
{
"default": 0.82,
"min": 0.0,
"max": 1.0,
"step": 0.01,
"tooltip": (
"Cosine-similarity threshold against canonical "
"refusal exemplars, above which an ambiguous "
"(short/hedge-y) response is treated as a "
"refusal. Only consulted when embedding_model "
"is set and the cheap lexical pass didn't "
"already catch it."
),
},
),
},
"optional": {"options": ("OLLAMA_OPTIONS",)},
}
RETURN_TYPES = ("OLLAMA_OPTIONS",)
RETURN_NAMES = ("options",)
FUNCTION = "set_refusal_retry"
CATEGORY = "dv/ollama/options"
def set_refusal_retry(self, enabled, embedding_model, threshold, options=None):
return (
_merge_option(
options,
"refusal_retry",
{
"enabled": enabled,
"embedding_model": embedding_model.strip(),
"threshold": threshold,
},
),
)
class OllamaOptionExtraBody:
@classmethod
def INPUT_TYPES(s):
+84
View File
@@ -558,3 +558,87 @@ def test_chat_text_only_content_stays_plain_string(monkeypatch):
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "hi"}]
# ---------------------------------------------------------------------------
# refusal-retry — parity with test_ollama_provider.py's coverage (ADR: this
# is a model-behavior concern, not a backend one — both providers wire the
# same comfydv._llm.retry.is_refusal() into their own retry loop).
# ---------------------------------------------------------------------------
def test_chat_retries_on_lexical_refusal_and_returns_clean_second_attempt(
monkeypatch,
):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
if len(calls) == 1:
return {"choices": [{"message": {"content": "I cannot generate that."}}]}
return {"choices": [{"message": {"content": "a real, on-topic answer"}}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
result = _run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b",
[Message(role="user", content="hi")],
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
)
)
assert result == "a real, on-topic answer"
assert len(calls) == 2
def test_chat_refusal_retry_disabled_returns_refusal_text_unchanged(monkeypatch):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"choices": [{"message": {"content": "I cannot generate that."}}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").chat(
"gemma-3-4b", [Message(role="user", content="hi")]
)
)
assert result == "I cannot generate that."
assert len(calls) == 1
def test_embed_returns_vector_from_v1_embeddings(monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
captured["payload"] = payload
return {"data": [{"embedding": [0.4, 0.5, 0.6], "index": 0}]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").embed("nomic-embed-text", "hello")
)
assert result == [0.4, 0.5, 0.6]
assert captured["url"] == "http://localhost:8080/v1/embeddings"
assert captured["payload"] == {"model": "nomic-embed-text", "input": "hello"}
def test_embed_returns_none_when_no_embedding_model_configured(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
raise RuntimeError("llama-server returned HTTP 404 for /v1/embeddings")
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
LlamaCppProvider("http://localhost:8080").embed("nomic-embed-text", "hello")
)
assert result is None
+106
View File
@@ -381,6 +381,112 @@ def test_chat_structured_retry_sleeps_between_attempts(monkeypatch):
assert sleep_calls == [chat_mod.RETRY_BACKOFF_SECS]
# ---------------------------------------------------------------------------
# refusal-retry — parity with test_ollama_provider.py's coverage (this
# module serves LlamaCppProvider.chat_structured() — see
# comfydv._llm.retry.is_refusal and OllamaOptionRefusalRetry).
# ---------------------------------------------------------------------------
def test_chat_structured_retries_on_refusal_and_returns_clean_second_attempt(
monkeypatch,
):
refused = _Widget(name="I cannot generate that content.", count=1)
clean = _Widget(name="clean", count=2)
fake = _FakeAgent([refused, clean])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
result = _run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=2,
)
)
assert result == clean
assert len(fake.calls) == 2
def test_chat_structured_refusal_retry_disabled_returns_refusal_unchanged(
monkeypatch,
):
refused = _Widget(name="I cannot generate that content.", count=1)
fake = _FakeAgent([refused])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
result = _run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
)
)
assert result == refused
assert len(fake.calls) == 1
def test_chat_structured_refusal_retry_exhausted_raises(monkeypatch):
refused = _Widget(name="I cannot help with this.", count=1)
fake = _FakeAgent([refused, refused])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
with pytest.raises(RuntimeError, match="failed validation"):
_run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=1,
)
)
def test_chat_structured_refusal_retry_uses_embed_fn(monkeypatch):
from comfydv._llm.retry import REFUSAL_EXEMPLARS
refused = _Widget(name="not today, sorry", count=1)
clean = _Widget(name="a clean value", count=2)
fake = _FakeAgent([refused, clean])
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
embed_calls = []
async def fake_embed(text):
embed_calls.append(text)
if "not today" in text or text in REFUSAL_EXEMPLARS:
return [1.0, 0.0]
return [0.0, 1.0]
result = _run_async(
chat_mod.chat_structured(
base_url="http://localhost:11434/v1",
model="llama3",
messages=_messages(),
schema=_Widget,
options={
"refusal_retry": {
"enabled": True,
"embedding_model": "nomic-embed-text",
"threshold": 0.5,
}
},
max_retries=2,
embed_fn=fake_embed,
)
)
assert result == clean
assert embed_calls # embedding path was actually exercised
def test_history_to_messages_preserves_order_and_roles():
from pydantic_ai.messages import ModelRequest, ModelResponse
+210 -2
View File
@@ -1,8 +1,24 @@
"""Tests for comfydv._llm.retry — shared retry-on-blank-output helpers used
by both providers' chat() and the shared chat_structured() helper.
by both providers' chat() and the shared chat_structured() helper, plus the
refusal/deflection detector that rides the same retry-with-a-new-seed
mechanism.
Refusal-detection tests here are pure-logic only — no provider/HTTP
involved. See test_ollama_provider.py, test_llamacpp_provider.py, and
test_llm_chat_structured.py for the retry-loop integration (does a detected
refusal actually trigger a reseeded retry).
"""
from comfydv._llm.retry import next_seed
import pytest
from comfydv._llm.ollama_provider import _run_async
from comfydv._llm.retry import (
cosine_similarity,
is_ambiguous,
is_lexical_refusal,
is_refusal,
next_seed,
)
def test_next_seed_attempt_one_is_zero_by_default():
@@ -22,3 +38,195 @@ def test_next_seed_starts_from_pinned_base():
def test_next_seed_ignores_non_int_seed():
assert next_seed({"seed": "not-an-int"}, 2) == 1
# ---------------------------------------------------------------------------
# Refusal/deflection detection
# ---------------------------------------------------------------------------
class TestLexicalDetection:
def test_blatant_refusal_openers_match(self):
blatant = [
"I cannot generate that image for you.",
"I'm not able to help with this request.",
"I don't feel comfortable creating this content.",
"I'm sorry, but I can't produce that description.",
"As an AI, I'm unable to depict this scenario.",
"I must decline this request.",
"This request violates my usage policies.",
"That's not something I'm able to help with right now.",
]
for text in blatant:
assert is_lexical_refusal(text), f"expected refusal match: {text!r}"
def test_ordinary_content_does_not_match(self):
ordinary = [
"The subject turns to face the camera and smiles warmly.",
"A person cannot simply walk into Mordor, the guide joked.",
"I can help you plan a birthday party for your dog.",
"",
]
for text in ordinary:
assert not is_lexical_refusal(text), f"unexpected match: {text!r}"
class TestAmbiguityHeuristic:
def test_short_response_is_ambiguous(self):
assert is_ambiguous("Sorry, can't do that one.")
def test_long_response_without_hedge_keywords_is_not_ambiguous(self):
long_text = "The subject rotates smoothly toward the lens. " * 20
assert len(long_text) >= 400
assert not is_ambiguous(long_text)
def test_long_response_with_hedge_keyword_is_ambiguous(self):
long_text = "Unfortunately, " + "this touches on a sensitive area. " * 20
assert len(long_text) >= 400
assert is_ambiguous(long_text)
def test_blank_text_is_not_ambiguous(self):
assert not is_ambiguous(" ")
class TestCosineSimilarity:
def test_identical_vectors_score_one(self):
assert cosine_similarity([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0)
def test_orthogonal_vectors_score_zero(self):
assert cosine_similarity([1.0, 0.0], [0.0, 1.0]) == pytest.approx(0.0)
def test_opposite_vectors_score_negative_one(self):
assert cosine_similarity([1.0, 0.0], [-1.0, 0.0]) == pytest.approx(-1.0)
def test_mismatched_lengths_return_zero(self):
assert cosine_similarity([1.0, 0.0], [1.0, 0.0, 0.0]) == 0.0
def test_empty_vectors_return_zero(self):
assert cosine_similarity([], []) == 0.0
class TestIsRefusalHybrid:
def test_blank_text_is_never_a_refusal(self):
assert not _run_async(is_refusal(""))
assert not _run_async(is_refusal(" "))
def test_lexical_match_short_circuits_without_embedding(self):
calls = []
async def embed_fn(text):
calls.append(text)
return [1.0, 0.0]
result = _run_async(
is_refusal("I cannot generate that image for you.", embed_fn=embed_fn)
)
assert result is True
assert calls == [] # never reached the embedding step
def test_long_clean_response_skips_embedding_and_is_not_refusal(self):
calls = []
async def embed_fn(text):
calls.append(text)
return [1.0, 0.0]
long_text = "The subject rotates smoothly toward the lens. " * 20
result = _run_async(is_refusal(long_text, embed_fn=embed_fn))
assert result is False
assert calls == [] # not ambiguous -> never reaches the embedding step
def test_no_embed_fn_degrades_to_lexical_only(self):
# Ambiguous (short), no lexical match, no embed_fn -> can't check further
assert not _run_async(is_refusal("Not today, sorry.", embed_fn=None))
def test_ambiguous_response_above_threshold_is_refusal(self):
async def embed_fn(text):
# Exemplars and a near-identical short "refusal-ish" probe get a
# high similarity score; distinguish by a marker substring.
if "PROBE" in text:
return [1.0, 0.0]
return [0.99, 0.14] # cos-sim with [1,0] is ~0.99
result = _run_async(
is_refusal(
"PROBE: not comfortable with this one",
embed_fn=embed_fn,
embed_cache_key="test-model",
threshold=0.8,
)
)
assert result is True
def test_ambiguous_response_below_threshold_is_not_refusal(self):
async def embed_fn(text):
if "PROBE" in text:
return [1.0, 0.0]
return [0.0, 1.0] # orthogonal -> cos-sim 0.0
result = _run_async(
is_refusal(
"PROBE: a short reply",
embed_fn=embed_fn,
embed_cache_key="test-model-2",
threshold=0.8,
)
)
assert result is False
def test_exemplar_embeddings_cached_across_calls(self):
exemplar_calls = {"n": 0}
async def embed_fn(text):
if "PROBE" in text:
return [1.0, 0.0]
exemplar_calls["n"] += 1
return [1.0, 0.0]
_run_async(
is_refusal(
"PROBE: first ambiguous call",
embed_fn=embed_fn,
embed_cache_key="cache-key-shared",
threshold=0.5,
)
)
first_count = exemplar_calls["n"]
assert first_count > 0
_run_async(
is_refusal(
"PROBE: second ambiguous call",
embed_fn=embed_fn,
embed_cache_key="cache-key-shared",
threshold=0.5,
)
)
# Exemplar embeddings reused from cache -> no additional exemplar calls
assert exemplar_calls["n"] == first_count
def test_embed_fn_failure_degrades_to_not_refused(self):
async def failing_embed_fn(text):
raise RuntimeError("server unreachable")
result = _run_async(
is_refusal(
"Not comfortable with this one, sorry.",
embed_fn=failing_embed_fn,
embed_cache_key="unreachable-model",
)
)
assert result is False
def test_embed_fn_returning_none_degrades_to_not_refused(self):
async def none_embed_fn(text):
return None
result = _run_async(
is_refusal(
"Not comfortable with this one, sorry.",
embed_fn=none_embed_fn,
embed_cache_key="no-embeddings-model",
)
)
assert result is False
+29
View File
@@ -50,6 +50,7 @@ from comfydv.ollama import (
OllamaOptionDisableThinking,
OllamaOptionExtraBody,
OllamaOptionMaxTokens,
OllamaOptionRefusalRetry,
OllamaOptionRepeatPenalty,
OllamaOptionSeed,
OllamaOptionTemperature,
@@ -732,6 +733,34 @@ class TestUS5ComposableOptions:
)
assert opts == {"temperature": 0.5, "think": False}
def test_refusal_retry_default_merges_config_dict(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True, embedding_model="", threshold=0.82
)
assert opts == {
"refusal_retry": {
"enabled": True,
"embedding_model": "",
"threshold": 0.82,
}
}
def test_refusal_retry_strips_embedding_model_whitespace(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=True, embedding_model=" nomic-embed-text ", threshold=0.9
)
assert opts["refusal_retry"]["embedding_model"] == "nomic-embed-text"
def test_refusal_retry_merges_existing_options(self):
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
enabled=False,
embedding_model="",
threshold=0.82,
options={"temperature": 0.5},
)
assert opts["temperature"] == 0.5
assert opts["refusal_retry"]["enabled"] is False
def test_extra_body_merges_json(self):
(opts,) = OllamaOptionExtraBody().set_extra_body(
extra_body_json='{"stop": ["</s>"]}', options={"temperature": 0.0}
+251
View File
@@ -19,6 +19,7 @@ import pytest
import comfydv._llm.ollama_provider as provider_mod
from comfydv._llm.ollama_provider import OllamaProvider, _run_async
from comfydv._llm.provider import Message, ModelStatus
from comfydv._llm.retry import REFUSAL_EXEMPLARS
@pytest.fixture(autouse=True)
@@ -924,3 +925,253 @@ def test_chat_text_only_payload_omits_images_key(monkeypatch):
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "hi"}]
# ---------------------------------------------------------------------------
# refusal-retry — comfydv-level "refusal_retry" options convention (see
# OllamaOptionRefusalRetry / comfydv._llm.retry.is_refusal)
# ---------------------------------------------------------------------------
def test_chat_retries_on_lexical_refusal_and_returns_clean_second_attempt(
monkeypatch,
):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
if len(calls) == 1:
return {"message": {"content": "I cannot generate that for you."}}
return {"message": {"content": "a real, on-topic answer"}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
result = _run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3",
[Message(role="user", content="hi")],
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
)
)
assert result == "a real, on-topic answer"
assert len(calls) == 2
# next_seed(): attempt 2 gets a bumped seed, not a repeat of attempt 1
assert calls[1]["options"]["seed"] == 1
def test_chat_refusal_retry_disabled_returns_refusal_text_unchanged(monkeypatch):
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"message": {"content": "I cannot generate that for you."}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", [Message(role="user", content="hi")]
)
)
assert result == "I cannot generate that for you."
assert len(calls) == 1
def test_chat_refused_response_is_never_cached(monkeypatch):
"""A refused attempt must not poison the cache the way a genuine success
does — otherwise every subsequent identical request would replay the
refusal forever (the exact bug the disable-thinking fix worked around
for a different failure mode)."""
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return {"message": {"content": "I cannot generate that for you."}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
messages = [Message(role="user", content="hi")]
options = {"refusal_retry": {"enabled": True, "embedding_model": ""}}
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", messages, options=options, max_retries=0
)
)
_run_async(
OllamaProvider("http://localhost:11434").chat(
"llama3", messages, options=options, max_retries=0
)
)
# Every attempt actually hit the network — nothing was ever cached.
assert len(calls) == 2
def test_chat_structured_retries_on_refusal_embedded_in_json_field(monkeypatch):
from pydantic import BaseModel
class Widget(BaseModel):
name: str
responses = [
{"message": {"content": '{"name": "I cannot generate that content."}'}},
{"message": {"content": '{"name": "a clean value"}'}},
]
calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return responses[len(calls) - 1]
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
result = _run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=2,
)
)
assert result == Widget(name="a clean value")
assert len(calls) == 2
def test_chat_structured_refusal_retry_exhausted_raises(monkeypatch):
from pydantic import BaseModel
class Widget(BaseModel):
name: str
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": '{"name": "I cannot help with this."}'}}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
with pytest.raises(RuntimeError, match="failed validation"):
_run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
max_retries=1,
)
)
def test_chat_structured_refusal_retry_uses_embedding_fallback(monkeypatch):
"""embedding_model set -> embed() is actually consulted for an ambiguous
(non-lexical-match) response, not just the free regex pass."""
from pydantic import BaseModel
class Widget(BaseModel):
name: str
responses = [
{"message": {"content": '{"name": "not today, sorry"}'}},
{"message": {"content": '{"name": "a clean value"}'}},
]
calls = []
embed_calls = []
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls.append(payload)
return responses[len(calls) - 1]
async def fake_embed(self, model, text):
embed_calls.append((model, text))
# Exemplars and the refusal-flavored response land in the same
# direction (high cosine sim); the clean response is orthogonal.
if "not today" in text or text in REFUSAL_EXEMPLARS:
return [1.0, 0.0]
return [0.0, 1.0]
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
monkeypatch.setattr(OllamaProvider, "embed", fake_embed)
result = _run_async(
OllamaProvider("http://localhost:11434").chat_structured(
"llama3",
[Message(role="user", content="hi")],
Widget,
options={
"refusal_retry": {
"enabled": True,
"embedding_model": "nomic-embed-text",
"threshold": 0.5,
}
},
max_retries=2,
)
)
assert result == Widget(name="a clean value")
assert len(calls) == 2
assert embed_calls # embedding path was actually exercised
assert all(model == "nomic-embed-text" for model, _ in embed_calls)
# ---------------------------------------------------------------------------
# embed()
# ---------------------------------------------------------------------------
def test_embed_returns_vector_from_api_embed(monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
captured["payload"] = payload
return {"embeddings": [[0.1, 0.2, 0.3]]}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result == [0.1, 0.2, 0.3]
assert captured["url"] == "http://localhost:11434/api/embed"
assert captured["payload"] == {"model": "nomic-embed-text", "input": "hello"}
def test_embed_returns_none_on_unreachable_server(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
raise RuntimeError("Cannot reach Ollama at ...")
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result is None
def test_embed_returns_none_for_malformed_response(monkeypatch):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"unexpected": "shape"}
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
result = _run_async(
OllamaProvider("http://localhost:11434").embed("nomic-embed-text", "hello")
)
assert result is None
def test_embed_returns_none_for_empty_model_or_text():
provider = OllamaProvider("http://localhost:11434")
assert _run_async(provider.embed("", "hello")) is None
assert _run_async(provider.embed("nomic-embed-text", "")) is None