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:
co-authored by
Claude Sonnet 5
parent
2724f535f5
commit
1f138e5991
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
...
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user