diff --git a/src/comfydv/__init__.py b/src/comfydv/__init__.py index 50849f2..8394a4e 100644 --- a/src/comfydv/__init__.py +++ b/src/comfydv/__init__.py @@ -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", diff --git a/src/comfydv/_llm/chat.py b/src/comfydv/_llm/chat.py index a38e91b..d9912ea 100644 --- a/src/comfydv/_llm/chat.py +++ b/src/comfydv/_llm/chat.py @@ -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) diff --git a/src/comfydv/_llm/llamacpp_provider.py b/src/comfydv/_llm/llamacpp_provider.py index 83ea5d8..75a5f25 100644 --- a/src/comfydv/_llm/llamacpp_provider.py +++ b/src/comfydv/_llm/llamacpp_provider.py @@ -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 diff --git a/src/comfydv/_llm/ollama_provider.py b/src/comfydv/_llm/ollama_provider.py index 604c072..6e9a7d5 100644 --- a/src/comfydv/_llm/ollama_provider.py +++ b/src/comfydv/_llm/ollama_provider.py @@ -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 diff --git a/src/comfydv/_llm/provider.py b/src/comfydv/_llm/provider.py index e29bcc7..fb19dff 100644 --- a/src/comfydv/_llm/provider.py +++ b/src/comfydv/_llm/provider.py @@ -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. + """ + ... diff --git a/src/comfydv/_llm/retry.py b/src/comfydv/_llm/retry.py index afab55b..18dd21c 100644 --- a/src/comfydv/_llm/retry.py +++ b/src/comfydv/_llm/retry.py @@ -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 diff --git a/src/comfydv/ollama.py b/src/comfydv/ollama.py index b9b1e43..8e201d2 100644 --- a/src/comfydv/ollama.py +++ b/src/comfydv/ollama.py @@ -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): diff --git a/tests/test_llamacpp_provider.py b/tests/test_llamacpp_provider.py index 26ef960..15253f6 100644 --- a/tests/test_llamacpp_provider.py +++ b/tests/test_llamacpp_provider.py @@ -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 diff --git a/tests/test_llm_chat_structured.py b/tests/test_llm_chat_structured.py index d1cfd46..df28e40 100644 --- a/tests/test_llm_chat_structured.py +++ b/tests/test_llm_chat_structured.py @@ -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 diff --git a/tests/test_llm_retry.py b/tests/test_llm_retry.py index 7b3ecf0..83a6f7d 100644 --- a/tests/test_llm_retry.py +++ b/tests/test_llm_retry.py @@ -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 diff --git a/tests/test_ollama.py b/tests/test_ollama.py index dada8f6..7972786 100644 --- a/tests/test_ollama.py +++ b/tests/test_ollama.py @@ -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": [""]}', options={"temperature": 0.0} diff --git a/tests/test_ollama_provider.py b/tests/test_ollama_provider.py index a232c7d..8ee3af0 100644 --- a/tests/test_ollama_provider.py +++ b/tests/test_ollama_provider.py @@ -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