Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f16da579bc | ||
|
|
7a40d22230 | ||
|
|
42c614d5f8 | ||
|
|
a0f391f212 | ||
|
|
3c5136af70 | ||
|
|
97249ee281 | ||
|
|
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",
|
||||
|
||||
@@ -28,6 +28,7 @@ pydantic-ai's internal one.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from typing import cast
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
@@ -46,7 +47,16 @@ 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,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
record_attempt_info,
|
||||
)
|
||||
|
||||
_STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS = (
|
||||
UnexpectedModelBehavior,
|
||||
@@ -126,6 +136,9 @@ async def chat_structured(
|
||||
options: dict | None = None,
|
||||
max_retries: int = 2,
|
||||
timeout_secs: float = 300.0,
|
||||
embed_fn: EmbedFn | None = None,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Call ``model`` at ``base_url`` (an OpenAI-compatible ``/v1`` root) and
|
||||
return a validated instance of ``schema``.
|
||||
@@ -155,19 +168,21 @@ 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(
|
||||
"chat_structured requires the last message to have role='user'"
|
||||
)
|
||||
|
||||
agent = _build_agent(
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
schema=schema,
|
||||
headers=headers,
|
||||
timeout_secs=timeout_secs,
|
||||
)
|
||||
history = _history_to_messages(messages)
|
||||
prompt = _user_prompt_content(messages[-1])
|
||||
think = None
|
||||
@@ -175,6 +190,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
|
||||
@@ -189,7 +209,33 @@ async def chat_structured(
|
||||
total_attempts = max(0, min(int(max_retries), 5)) + 1
|
||||
last_error: Exception | None = None
|
||||
last_invalid_text = ""
|
||||
refusal_count = 0
|
||||
attempt_seed = (options or {}).get("seed", 0) if isinstance(options, dict) else 0
|
||||
attempt_timeout = timeout_secs
|
||||
|
||||
def _emit_retry_status(reason: str, attempt: int) -> None:
|
||||
if on_status is None or attempt >= total_attempts:
|
||||
return
|
||||
upcoming_seed = next_seed(options, attempt + 1)
|
||||
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||
on_status(
|
||||
format_retry_status(
|
||||
reason, attempt, total_attempts, upcoming_seed, upcoming_timeout
|
||||
)
|
||||
)
|
||||
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||
# Rebuilt each attempt so the escalated timeout actually takes
|
||||
# effect — httpx.AsyncClient's timeout is fixed at construction,
|
||||
# not mutable per-request.
|
||||
agent = _build_agent(
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
schema=schema,
|
||||
headers=headers,
|
||||
timeout_secs=attempt_timeout,
|
||||
)
|
||||
attempt_settings = dict(model_settings) if model_settings else {}
|
||||
if attempt > 1:
|
||||
# Confirmed live: a freshly-loaded model's first structured-output
|
||||
@@ -202,6 +248,7 @@ async def chat_structured(
|
||||
# llama-server's OpenAI-compatible endpoints) and give it a beat
|
||||
# via RETRY_BACKOFF_SECS in case it's still finishing loading.
|
||||
seed = next_seed(options, attempt)
|
||||
attempt_seed = seed
|
||||
attempt_settings["seed"] = seed
|
||||
if "extra_body" in attempt_settings:
|
||||
# beacon-reviewer caught this: if a caller pinned options["seed"],
|
||||
@@ -233,13 +280,54 @@ 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),
|
||||
custom_phrases=tuple(refusal_cfg.get("custom_phrases") or ()),
|
||||
)
|
||||
if refused:
|
||||
refusal_count += 1
|
||||
last_error = RuntimeError("refusal/deflection detected")
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||
)
|
||||
return output
|
||||
except _STRUCTURED_OUTPUT_FAILURE_EXCEPTIONS as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = str(exc)
|
||||
_emit_retry_status("Structured output failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=total_attempts,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
raise RuntimeError(
|
||||
f"chat_structured: response failed validation against schema after "
|
||||
f"{total_attempts} attempt(s) (model={model!r}). Last error: "
|
||||
|
||||
@@ -15,12 +15,28 @@ exist otherwise (spec.md FR-006).
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
|
||||
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,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
record_attempt_info,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -183,13 +199,27 @@ class LlamaCppProvider:
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> 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
|
||||
custom_phrases: tuple[str, ...] = ()
|
||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
|
||||
if 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 = ""
|
||||
refusal_count = 0
|
||||
attempt_seed = 0
|
||||
attempt_timeout = timeout_secs
|
||||
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
"messages": payload_messages,
|
||||
@@ -218,6 +248,13 @@ class LlamaCppProvider:
|
||||
# spec's actual top-level "seed" field, so it takes effect
|
||||
# against llama-server's /v1/chat/completions.
|
||||
payload["seed"] = next_seed(options, attempt)
|
||||
attempt_seed = payload["seed"]
|
||||
else:
|
||||
# attempt 1 never sets the top-level "seed" field above (only
|
||||
# retries do) — fall back to whatever the caller pinned in
|
||||
# options, so attempt_info/seed_used reports the real seed in
|
||||
# play even on a first-attempt success, not a stale 0.
|
||||
attempt_seed = (options or {}).get("seed", 0)
|
||||
|
||||
cache_key = _cache_key(
|
||||
"llamacpp_chat",
|
||||
@@ -231,12 +268,19 @@ class LlamaCppProvider:
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
return cached
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/v1/chat/completions",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
timeout=attempt_timeout,
|
||||
headers=self.headers,
|
||||
)
|
||||
choices = result.get("choices") or []
|
||||
@@ -245,13 +289,60 @@ class LlamaCppProvider:
|
||||
if choices
|
||||
else ""
|
||||
)
|
||||
retry_reason: str | None = None
|
||||
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),
|
||||
custom_phrases=custom_phrases,
|
||||
)
|
||||
if not refused:
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(
|
||||
attempt, total_attempts, attempt_seed
|
||||
)
|
||||
)
|
||||
return response_text
|
||||
refusal_count += 1
|
||||
retry_reason = "Refusal/deflection detected"
|
||||
else:
|
||||
retry_reason = "Blank response"
|
||||
|
||||
if attempt < total_attempts:
|
||||
if on_status is not None:
|
||||
upcoming_seed = next_seed(options, attempt + 1)
|
||||
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||
on_status(
|
||||
format_retry_status(
|
||||
retry_reason,
|
||||
attempt,
|
||||
total_attempts,
|
||||
upcoming_seed,
|
||||
upcoming_timeout,
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=total_attempts,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
# Every attempt came back blank — never raises here (chat() has
|
||||
# never validated its output, unlike chat_structured()); return the
|
||||
# last (blank) attempt uncached so the next queue run tries fresh.
|
||||
@@ -265,6 +356,8 @@ class LlamaCppProvider:
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
from .chat import chat_structured as _chat_structured_impl
|
||||
|
||||
@@ -280,8 +373,25 @@ class LlamaCppProvider:
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=(options or {}).get("seed", 0),
|
||||
attempts=1,
|
||||
timeout_secs=timeout_secs,
|
||||
refusals=0,
|
||||
)
|
||||
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 +401,40 @@ class LlamaCppProvider:
|
||||
options=options,
|
||||
max_retries=max_retries,
|
||||
timeout_secs=timeout_secs,
|
||||
embed_fn=embed_fn,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
_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
|
||||
|
||||
@@ -15,11 +15,20 @@ import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
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,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
record_attempt_info,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -102,6 +111,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)
|
||||
@@ -343,6 +368,8 @@ class OllamaProvider:
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
if any(m.images for m in messages):
|
||||
await _require_vision_capability(self.host, model, self.headers)
|
||||
@@ -353,14 +380,27 @@ 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
|
||||
custom_phrases: tuple[str, ...] = ()
|
||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
|
||||
if 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
|
||||
refusal_count = 0
|
||||
attempt_seed = 0
|
||||
attempt_timeout = timeout_secs
|
||||
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
attempt_options = dict(options) if options else {}
|
||||
if attempt > 1:
|
||||
attempt_options["seed"] = next_seed(options, attempt)
|
||||
attempt_seed = attempt_options.get("seed", 0)
|
||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
@@ -385,18 +425,55 @@ class OllamaProvider:
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
return cached
|
||||
|
||||
result = await _post_json(
|
||||
f"{self.host}/api/chat",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
timeout=attempt_timeout,
|
||||
headers=self.headers,
|
||||
)
|
||||
response_text = result.get("message", {}).get("content", "")
|
||||
retry_reason: str | None = None
|
||||
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),
|
||||
custom_phrases=custom_phrases,
|
||||
)
|
||||
if not refused:
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(
|
||||
attempt, total_attempts, attempt_seed
|
||||
)
|
||||
)
|
||||
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.
|
||||
refusal_count += 1
|
||||
retry_reason = "Refusal/deflection detected"
|
||||
|
||||
# done: false alongside blank content is a distinct signal from
|
||||
# an ordinary blank generation — it's Ollama answering before
|
||||
@@ -406,10 +483,34 @@ class OllamaProvider:
|
||||
# raised on below instead of silently returned like a real
|
||||
# blank generation would be.
|
||||
incomplete = result.get("done") is False
|
||||
if retry_reason is None and not response_text.strip():
|
||||
retry_reason = (
|
||||
"Model still loading/swapping" if incomplete else "Blank response"
|
||||
)
|
||||
|
||||
if attempt < total_attempts:
|
||||
if on_status is not None and retry_reason is not None:
|
||||
upcoming_seed = next_seed(options, attempt + 1)
|
||||
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||
on_status(
|
||||
format_retry_status(
|
||||
retry_reason,
|
||||
attempt,
|
||||
total_attempts,
|
||||
upcoming_seed,
|
||||
upcoming_timeout,
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=total_attempts,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
|
||||
if incomplete:
|
||||
raise RuntimeError(
|
||||
f"Ollama returned an incomplete response after "
|
||||
@@ -432,6 +533,8 @@ class OllamaProvider:
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Native ``/api/chat`` + ``"format"`` (grammar-constrained JSON
|
||||
decoding), not the shared pydantic-ai ``chat.py`` helper.
|
||||
@@ -457,6 +560,14 @@ 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
|
||||
custom_phrases: tuple[str, ...] = ()
|
||||
if refusal_cfg and refusal_cfg.get("enabled"):
|
||||
custom_phrases = tuple(refusal_cfg.get("custom_phrases") or ())
|
||||
if 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,
|
||||
@@ -469,16 +580,39 @@ class OllamaProvider:
|
||||
)
|
||||
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
|
||||
if hit:
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=(options or {}).get("seed", 0),
|
||||
attempts=1,
|
||||
timeout_secs=timeout_secs,
|
||||
refusals=0,
|
||||
)
|
||||
return schema.model_validate(cached)
|
||||
|
||||
total_attempts = max(0, min(int(max_retries), 5)) + 1
|
||||
last_error: Exception | None = None
|
||||
last_invalid_text = ""
|
||||
refusal_count = 0
|
||||
attempt_seed = 0
|
||||
attempt_timeout = timeout_secs
|
||||
|
||||
def _emit_retry_status(reason: str, attempt: int) -> None:
|
||||
if on_status is None or attempt >= total_attempts:
|
||||
return
|
||||
upcoming_seed = next_seed(options, attempt + 1)
|
||||
upcoming_timeout = next_timeout_secs(timeout_secs, attempt + 1)
|
||||
on_status(
|
||||
format_retry_status(
|
||||
reason, attempt, total_attempts, upcoming_seed, upcoming_timeout
|
||||
)
|
||||
)
|
||||
|
||||
for attempt in range(1, total_attempts + 1):
|
||||
attempt_options = dict(options) if options else {}
|
||||
if attempt > 1:
|
||||
attempt_options["seed"] = next_seed(options, attempt)
|
||||
attempt_seed = attempt_options.get("seed", 0)
|
||||
attempt_timeout = next_timeout_secs(timeout_secs, attempt)
|
||||
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
@@ -497,12 +631,13 @@ class OllamaProvider:
|
||||
result = await _post_json(
|
||||
f"{self.host}/api/chat",
|
||||
payload,
|
||||
timeout=timeout_secs,
|
||||
timeout=attempt_timeout,
|
||||
headers=self.headers,
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = str(exc)
|
||||
_emit_retry_status("Request failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
@@ -513,15 +648,84 @@ class OllamaProvider:
|
||||
except ValidationError as exc:
|
||||
last_error = exc
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Schema validation failed", attempt)
|
||||
if attempt < total_attempts:
|
||||
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),
|
||||
custom_phrases=custom_phrases,
|
||||
)
|
||||
if refused:
|
||||
refusal_count += 1
|
||||
last_error = RuntimeError("refusal/deflection detected")
|
||||
last_invalid_text = content
|
||||
_emit_retry_status("Refusal/deflection detected", attempt)
|
||||
if attempt < total_attempts:
|
||||
await asyncio.sleep(RETRY_BACKOFF_SECS)
|
||||
continue
|
||||
|
||||
_CHAT_RESPONSE_CACHE.set(cache_key, parsed.model_dump())
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=attempt,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
if on_status is not None and attempt > 1:
|
||||
on_status(
|
||||
format_recovered_status(attempt, total_attempts, attempt_seed)
|
||||
)
|
||||
return parsed
|
||||
|
||||
record_attempt_info(
|
||||
attempt_info,
|
||||
seed=attempt_seed,
|
||||
attempts=total_attempts,
|
||||
timeout_secs=attempt_timeout,
|
||||
refusals=refusal_count,
|
||||
)
|
||||
raise RuntimeError(
|
||||
f"chat_structured: response failed validation against schema after "
|
||||
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
|
||||
|
||||
@@ -8,6 +8,7 @@ project-management/ADRs/ADR-007-llm-provider-adapter-pattern.md and
|
||||
specs/007-llm-provider-abstraction/contracts/llm_provider_protocol.md.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Literal, Protocol
|
||||
|
||||
@@ -83,6 +84,8 @@ class LLMProvider(Protocol):
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> str:
|
||||
"""Free-text chat response.
|
||||
|
||||
@@ -92,7 +95,9 @@ class LLMProvider(Protocol):
|
||||
into normal behavior. Still returns the (possibly blank) last
|
||||
attempt's text rather than raising if every retry comes back blank —
|
||||
this method has never validated its output, unlike
|
||||
``chat_structured()``.
|
||||
``chat_structured()``. Each retry's request timeout also escalates
|
||||
(``_llm/retry.py``'s ``next_timeout_secs``) rather than reusing the
|
||||
same budget that just ran out.
|
||||
|
||||
ADR-010: ``options`` may carry a ``"think"`` key (bool) to disable a
|
||||
"thinking"-capable model's chain-of-thought reasoning — every
|
||||
@@ -101,6 +106,18 @@ class LLMProvider(Protocol):
|
||||
``chat_template_kwargs``/``reasoning_effort`` in the request body),
|
||||
since neither backend recognizes a literal ``"think"`` key nested
|
||||
inside a generic options object.
|
||||
|
||||
``attempt_info``, if given, is populated in place with the retry
|
||||
loop's final outcome (seed/timeout used, attempt count, refusal
|
||||
count) via ``_llm/retry.py``'s ``record_attempt_info`` — an optional
|
||||
out-param, not a return-type change, so existing callers that don't
|
||||
pass it see no behavior change.
|
||||
|
||||
``on_status``, if given, is called synchronously at each retry
|
||||
boundary with a one-line human-readable status (see
|
||||
``_llm/retry.py``'s ``format_retry_status``/``format_recovered_status``)
|
||||
— a live counterpart to ``attempt_info``, which only reports the
|
||||
final outcome after the call returns.
|
||||
"""
|
||||
...
|
||||
|
||||
@@ -112,6 +129,8 @@ class LLMProvider(Protocol):
|
||||
options: dict | None = None,
|
||||
timeout_secs: float = 300.0,
|
||||
max_retries: int = 2,
|
||||
attempt_info: dict | None = None,
|
||||
on_status: Callable[[str], None] | None = None,
|
||||
) -> BaseModel:
|
||||
"""Schema-validated chat response.
|
||||
|
||||
@@ -121,6 +140,22 @@ class LLMProvider(Protocol):
|
||||
field.
|
||||
|
||||
ADR-010: see ``chat()`` — same ``options["think"]`` convention,
|
||||
same per-provider translation.
|
||||
same per-provider translation, same escalating per-attempt timeout,
|
||||
and the same ``attempt_info``/``on_status`` conventions.
|
||||
"""
|
||||
...
|
||||
|
||||
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,289 @@ def next_seed(options: dict | None, attempt: int) -> int:
|
||||
if options and isinstance(options.get("seed"), int):
|
||||
base = options["seed"]
|
||||
return base + (attempt - 1)
|
||||
|
||||
|
||||
def next_timeout_secs(base_timeout: float, attempt: int) -> float:
|
||||
"""Escalating per-attempt timeout for retries (1-indexed ``attempt``).
|
||||
|
||||
Attempt 1 gets the caller's own ``timeout_secs`` unchanged; each retry
|
||||
multiplies it by the attempt number. A request that timed out may
|
||||
genuinely need more time — a slow-to-load or heavily-loaded model, a
|
||||
large prompt — not just an identical retry under the same budget it
|
||||
just failed to meet.
|
||||
"""
|
||||
return base_timeout * attempt
|
||||
|
||||
|
||||
def record_attempt_info(
|
||||
attempt_info: dict | None,
|
||||
*,
|
||||
seed: int,
|
||||
attempts: int,
|
||||
timeout_secs: float,
|
||||
refusals: int,
|
||||
) -> None:
|
||||
"""Populate an optional caller-supplied dict with the retry loop's
|
||||
final outcome — the seed/timeout actually used, how many attempts it
|
||||
took, and how many were refusal-triggered.
|
||||
|
||||
A plain out-param rather than a return-type change, so it's fully
|
||||
backward compatible: a caller that doesn't pass ``attempt_info`` sees
|
||||
no change in behavior at all. ``ChatCompletion`` uses this to expose
|
||||
the seed actually used as a node output and to build a UI status line
|
||||
when a retry/refusal happened.
|
||||
"""
|
||||
if attempt_info is None:
|
||||
return
|
||||
attempt_info.update(
|
||||
{
|
||||
"seed": seed,
|
||||
"attempts": attempts,
|
||||
"timeout_secs": timeout_secs,
|
||||
"refusals": refusals,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
OnStatus = Callable[[str], None]
|
||||
"""A caller-supplied, synchronous, best-effort progress callback — see
|
||||
``format_retry_status``/``format_recovered_status``. Not async: providers
|
||||
call it inline mid-retry-loop, and the one real implementation
|
||||
(``ChatCompletion``'s closure over ``PromptServer.send_progress_text``) is
|
||||
itself synchronous, so there's nothing to await."""
|
||||
|
||||
|
||||
def format_retry_status(
|
||||
reason: str, attempt: int, total_attempts: int, seed: int, timeout_secs: float
|
||||
) -> str:
|
||||
"""One-line, human-readable status for ``on_status()`` callers — shown
|
||||
live on the node via ComfyUI's ``PromptServer.send_progress_text``
|
||||
(see ``ChatCompletion.chat()``). Centralized so every provider's retry
|
||||
loop describes a retry the same way rather than each inventing its own
|
||||
wording.
|
||||
"""
|
||||
return (
|
||||
f"⚠ {reason} on attempt {attempt}/{total_attempts} — "
|
||||
f"retrying with seed={seed}, timeout={timeout_secs:.0f}s"
|
||||
)
|
||||
|
||||
|
||||
def format_recovered_status(attempt: int, total_attempts: int, seed: int) -> str:
|
||||
"""Final status shown once a retry loop succeeds after >1 attempt —
|
||||
lets a live status left over from ``format_retry_status`` resolve to
|
||||
something other than a stale "retrying..." message."""
|
||||
return f"✅ Recovered on attempt {attempt}/{total_attempts} (seed={seed})"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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"\b(?:I\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|\s+am) (?:(?:not able|unable) to|restricted from)\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",
|
||||
"inappropriate",
|
||||
"instead, i",
|
||||
"i'd rather",
|
||||
"i would rather",
|
||||
"controversial",
|
||||
"harmful",
|
||||
"offensive",
|
||||
"disturbing",
|
||||
"explicit content",
|
||||
"nsfw",
|
||||
"consider an alternative",
|
||||
)
|
||||
|
||||
_AMBIGUOUS_LENGTH_THRESHOLD = 600
|
||||
"""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.
|
||||
|
||||
Note: ``is_refusal`` no longer applies this gate when the caller supplies
|
||||
``embed_fn`` (see below) — an explicitly configured embedding model always
|
||||
runs the check regardless of length/keywords. ``is_ambiguous`` is kept as a
|
||||
standalone, independently tested heuristic for other callers (e.g. logging
|
||||
"why was this worth embedding-checking") rather than wired into the hybrid
|
||||
detector itself."""
|
||||
|
||||
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, exemplars: tuple[str, ...] = REFUSAL_EXEMPLARS
|
||||
) -> list[list[float]]:
|
||||
"""Embed ``exemplars`` once per ``cache_key`` and reuse — the exemplar
|
||||
set only changes if the caller's custom phrases change (folded into
|
||||
``cache_key`` by the caller), or 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 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
|
||||
|
||||
|
||||
def _matches_custom_phrase(text: str, custom_phrases: tuple[str, ...]) -> bool:
|
||||
"""Case-insensitive substring match against user-supplied phrases."""
|
||||
if not custom_phrases:
|
||||
return False
|
||||
lowered = text.lower()
|
||||
return any(phrase.lower() in lowered for phrase in custom_phrases)
|
||||
|
||||
|
||||
async def is_refusal(
|
||||
text: str,
|
||||
*,
|
||||
embed_fn: EmbedFn | None = None,
|
||||
embed_cache_key: str = "",
|
||||
threshold: float = 0.82,
|
||||
custom_phrases: tuple[str, ...] = (),
|
||||
) -> bool:
|
||||
"""Hybrid refusal/deflection detector: free lexical pass first, then an
|
||||
embedding-similarity fallback whenever the caller has configured one.
|
||||
|
||||
``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; ``embed_fn`` present
|
||||
means the caller already opted in to the extra cost, so every non-blank,
|
||||
non-lexically-caught response gets checked — no further length/keyword
|
||||
gating. 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.
|
||||
|
||||
``custom_phrases`` lets a caller extend detection at runtime — e.g. a
|
||||
ComfyUI node field the user edits directly — without touching the
|
||||
shipped patterns/exemplars. Each phrase is checked two ways: a free
|
||||
case-insensitive substring match (same cost tier as the lexical pass,
|
||||
so it runs even with no ``embed_fn`` configured), and, when ``embed_fn``
|
||||
is present, folded in as additional exemplars for the similarity check
|
||||
so near-matches (not just exact substrings) of the user's phrases count
|
||||
too.
|
||||
"""
|
||||
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
|
||||
custom_phrases = tuple(p.strip() for p in custom_phrases if p and p.strip())
|
||||
if _matches_custom_phrase(text, custom_phrases):
|
||||
return True
|
||||
if embed_fn is None:
|
||||
return False
|
||||
# embed_fn only exists when the caller explicitly configured an
|
||||
# embedding_model — that's an opt-in to pay for the check, so run it on
|
||||
# every non-blank, non-lexically-caught response rather than gating
|
||||
# further on is_ambiguous. The length/keyword heuristic exists to avoid
|
||||
# *unwanted* embedding calls when no embedding model is configured (see
|
||||
# the embed_fn is None branch above); it has no reason to also suppress
|
||||
# calls once the caller has already asked for them, and doing so was
|
||||
# exactly what let the subtle/on-topic-looking deflections this feature
|
||||
# targets slip through undetected.
|
||||
exemplars = (
|
||||
REFUSAL_EXEMPLARS + custom_phrases if custom_phrases else REFUSAL_EXEMPLARS
|
||||
)
|
||||
exemplar_cache_key = (
|
||||
f"{embed_cache_key}|custom:{','.join(custom_phrases)}"
|
||||
if custom_phrases
|
||||
else embed_cache_key
|
||||
)
|
||||
try:
|
||||
exemplar_vecs = await _exemplar_embeddings(
|
||||
embed_fn, exemplar_cache_key, exemplars
|
||||
)
|
||||
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
|
||||
|
||||
+184
-6
@@ -509,8 +509,20 @@ def _coerce_structured_value(value, comfy_type: str):
|
||||
class ChatCompletion:
|
||||
OUTPUT_NODE = True
|
||||
|
||||
_BASE_RETURN_TYPES = ("STRING", "OLLAMA_HISTORY", "STRING")
|
||||
_BASE_RETURN_NAMES = ("response", "updated_history", "model_name")
|
||||
# "seed_used" is always the LAST output, after any dynamic structured-
|
||||
# output fields — not inserted right after model_name — so a schema's
|
||||
# fields keep starting at the same fixed index (3) they occupied before
|
||||
# this output existed. Already-wired workflows (e.g.
|
||||
# workflows/ltx-i2v-pipeline.json) link FormatString inputs to a
|
||||
# ChatCompletion node's dynamic field by output *index*; inserting a
|
||||
# new fixed output ahead of those fields would silently repoint every
|
||||
# such link at the wrong socket.
|
||||
_FIXED_RETURN_TYPES = ("STRING", "OLLAMA_HISTORY", "STRING")
|
||||
_FIXED_RETURN_NAMES = ("response", "updated_history", "model_name")
|
||||
_SEED_RETURN_TYPE = ("INT",)
|
||||
_SEED_RETURN_NAME = ("seed_used",)
|
||||
_BASE_RETURN_TYPES = _FIXED_RETURN_TYPES + _SEED_RETURN_TYPE
|
||||
_BASE_RETURN_NAMES = _FIXED_RETURN_NAMES + _SEED_RETURN_NAME
|
||||
|
||||
# Per-node-instance structured-output config, keyed by unique_id — same
|
||||
# pattern as FormatString.node_configs.
|
||||
@@ -579,8 +591,14 @@ class ChatCompletion:
|
||||
cls.RETURN_NAMES = cls._BASE_RETURN_NAMES
|
||||
return
|
||||
names = tuple(schema["properties"].keys())
|
||||
cls.RETURN_TYPES = cls._BASE_RETURN_TYPES + _comfy_types_for_schema(schema)
|
||||
cls.RETURN_NAMES = cls._BASE_RETURN_NAMES + names
|
||||
# Dynamic fields land between the fixed base outputs and the
|
||||
# trailing seed_used — see _SEED_RETURN_TYPE's comment above.
|
||||
cls.RETURN_TYPES = (
|
||||
cls._FIXED_RETURN_TYPES
|
||||
+ _comfy_types_for_schema(schema)
|
||||
+ cls._SEED_RETURN_TYPE
|
||||
)
|
||||
cls.RETURN_NAMES = cls._FIXED_RETURN_NAMES + names + cls._SEED_RETURN_NAME
|
||||
|
||||
def chat(
|
||||
self,
|
||||
@@ -624,6 +642,33 @@ class ChatCompletion:
|
||||
if user_images:
|
||||
messages[-1].images = user_images
|
||||
llm_options = dict(options) if options else None
|
||||
# Populated in place by the provider (see _llm/retry.py's
|
||||
# record_attempt_info) with the seed/timeout actually used and how
|
||||
# many attempts/refusals it took — an out-param rather than a
|
||||
# return-type change, so it works the same regardless of which
|
||||
# concrete provider ``client`` is.
|
||||
attempt_info: dict = {}
|
||||
|
||||
# Live counterpart to attempt_info: ComfyUI's own send_progress_text
|
||||
# mechanism (already used by core nodes like PreviewAny/gaussian
|
||||
# splat count) shows this text on the node WHILE it's still
|
||||
# executing, via a "progressText" widget the frontend creates
|
||||
# automatically — no custom JS needed on our side. Best-effort:
|
||||
# a failure here must never take down the actual chat call.
|
||||
on_status = None
|
||||
if unique_id and "comfy" in sys.modules:
|
||||
|
||||
def on_status(message: str) -> None:
|
||||
try:
|
||||
from server import PromptServer
|
||||
|
||||
PromptServer.instance.send_progress_text(message, unique_id)
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"Failed to send live retry status for node %s",
|
||||
unique_id,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
# Provider owns transport, caching, and — for structured_output — the
|
||||
# tool-calling/retry/validation mechanism (pydantic-ai, ADR-007).
|
||||
@@ -637,6 +682,8 @@ class ChatCompletion:
|
||||
llm_options,
|
||||
timeout_secs=float(timeout_secs),
|
||||
max_retries=max_retries,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -651,20 +698,44 @@ class ChatCompletion:
|
||||
llm_options,
|
||||
timeout_secs=float(timeout_secs),
|
||||
max_retries=max_retries,
|
||||
attempt_info=attempt_info,
|
||||
on_status=on_status,
|
||||
)
|
||||
)
|
||||
response_text = parsed.model_dump_json()
|
||||
|
||||
seed_used = attempt_info.get("seed", 0)
|
||||
attempts_made = attempt_info.get("attempts", 1)
|
||||
refusals = attempt_info.get("refusals", 0)
|
||||
|
||||
updated = list(history)
|
||||
updated.append({"role": "user", "content": prompt})
|
||||
updated.append({"role": "assistant", "content": response_text})
|
||||
n = len(updated)
|
||||
|
||||
# Visual feedback (US: "show me when a retry/refusal happened") —
|
||||
# surfaced in the node's existing text preview rather than a new UI
|
||||
# surface, so it's visible without any frontend/JS changes.
|
||||
status_line = ""
|
||||
if attempts_made > 1:
|
||||
status_line = (
|
||||
f"⚠️ Refusal/deflection detected — retried {refusals} "
|
||||
f"time(s), succeeded on attempt {attempts_made} with "
|
||||
f"seed={seed_used}.\n\n"
|
||||
if refusals
|
||||
else f"⚠️ Retried (blank/invalid response) — succeeded on "
|
||||
f"attempt {attempts_made} with seed={seed_used}.\n\n"
|
||||
)
|
||||
|
||||
ui_text = (
|
||||
f"{response_text}\n\n── History: {n} message(s) ──\n{_history_preview(updated)}"
|
||||
f"{status_line}{response_text}\n\n"
|
||||
f"── History: {n} message(s) ──\n{_history_preview(updated)}"
|
||||
if n > 2
|
||||
else response_text
|
||||
else f"{status_line}{response_text}"
|
||||
)
|
||||
|
||||
# seed_used is appended last, after any dynamic structured fields —
|
||||
# see ChatCompletion's _SEED_RETURN_TYPE comment for why.
|
||||
result_tuple = (response_text, updated, effective_model)
|
||||
if structured_output:
|
||||
assert schema is not None # structured_output implies this was parsed
|
||||
@@ -674,6 +745,7 @@ class ChatCompletion:
|
||||
for name, ctype in zip(schema["properties"].keys(), comfy_types)
|
||||
)
|
||||
result_tuple += extra
|
||||
result_tuple += (seed_used,)
|
||||
|
||||
return {
|
||||
"ui": {"text": [ui_text]},
|
||||
@@ -914,6 +986,112 @@ 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.
|
||||
|
||||
``custom_phrases`` (comma-separated) lets you add your own trigger
|
||||
phrases at runtime, without a code change/release — useful for a new
|
||||
deflection phrasing a specific model uses that the shipped patterns in
|
||||
``REFUSAL_LEXICAL_PATTERNS``/``REFUSAL_EXEMPLARS`` don't cover yet.
|
||||
"""
|
||||
|
||||
@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."
|
||||
),
|
||||
},
|
||||
),
|
||||
"custom_phrases": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Comma-separated phrases you want treated as "
|
||||
'refusals too, e.g. "I am restricted from, as an '
|
||||
'AI model, I must avoid". Checked as free, exact '
|
||||
"case-insensitive substrings (no embedding model "
|
||||
"needed) and, when embedding_model is set, also "
|
||||
"folded in as extra exemplars for the similarity "
|
||||
"check — lets you extend detection at runtime "
|
||||
"without waiting on a shipped pattern update."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"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, custom_phrases="", options=None
|
||||
):
|
||||
phrases = tuple(p.strip() for p in custom_phrases.split(",") if p and p.strip())
|
||||
return (
|
||||
_merge_option(
|
||||
options,
|
||||
"refusal_retry",
|
||||
{
|
||||
"enabled": enabled,
|
||||
"embedding_model": embedding_model.strip(),
|
||||
"threshold": threshold,
|
||||
"custom_phrases": phrases,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class OllamaOptionExtraBody:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import sys
|
||||
@@ -7,6 +8,28 @@ from .utils import any_type
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _preview_text(value) -> str:
|
||||
"""Best-effort text preview for RandomChoice's arbitrary-typed output.
|
||||
|
||||
Mirrors ComfyUI core's own ``PreviewAny`` node's value handling (str/
|
||||
number passthrough, else JSON, else ``str()``) rather than inventing a
|
||||
new convention — RandomChoice's output can be anything (an IMAGE
|
||||
tensor, a LATENT, a plain string), so this only needs to be "good
|
||||
enough to glance at," not a faithful repr of every type.
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, (int, float, bool)):
|
||||
return str(value)
|
||||
try:
|
||||
return json.dumps(value, default=str, indent=2)
|
||||
except Exception:
|
||||
try:
|
||||
return str(value)
|
||||
except Exception:
|
||||
return "<value could not be serialized>"
|
||||
|
||||
|
||||
class RandomChoice:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -25,15 +48,20 @@ class RandomChoice:
|
||||
|
||||
FUNCTION = "random_choice"
|
||||
|
||||
OUTPUT_NODE = False
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "dv/utils"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, **kwargs):
|
||||
return s.random_choice(s, **kwargs)
|
||||
# Unchanged from before the UI-preview addition: returns the raw
|
||||
# picked value (not the ui-wrapped dict random_choice() now returns)
|
||||
# so ComfyUI's change-detection comparison keeps working exactly as
|
||||
# it did previously.
|
||||
return s._pick(**kwargs)
|
||||
|
||||
def random_choice(self, **kwargs):
|
||||
@staticmethod
|
||||
def _pick(**kwargs):
|
||||
(
|
||||
random.seed(kwargs.get("seed"))
|
||||
if kwargs.get("seed")
|
||||
@@ -41,10 +69,13 @@ class RandomChoice:
|
||||
)
|
||||
input = [i for i in kwargs.items() if i[0] != "seed"]
|
||||
logger.debug("RandomChoice inputs: %s", input)
|
||||
return random.choice(input)[1]
|
||||
|
||||
def random_choice(self, **kwargs):
|
||||
try:
|
||||
choice = random.choice(input)[1]
|
||||
choice = self._pick(**kwargs)
|
||||
logger.debug("RandomChoice chose: %s", choice)
|
||||
return (choice,)
|
||||
return {"ui": {"text": [_preview_text(choice)]}, "result": (choice,)}
|
||||
except Exception as e:
|
||||
logger.error("RandomChoice: unexpected error: %s", e)
|
||||
raise
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
/**
|
||||
* preview_text.js — read-only output preview for comfydv's OUTPUT_NODE=True
|
||||
* nodes that return a ComfyUI "ui": {"text": [...]} payload
|
||||
* (ChatCompletion, FormatString, RandomChoice).
|
||||
*
|
||||
* ComfyUI does NOT auto-render an arbitrary node's ui.text — each node type
|
||||
* that wants one implements its own onExecuted handler. This mirrors core's
|
||||
* own ``PreviewAny`` node (comfy_extras/nodes_preview_any.py +
|
||||
* "Comfy.PreviewAny" in the frontend bundle) minus its Markdown/Plaintext
|
||||
* toggle, which none of these three nodes need.
|
||||
*/
|
||||
|
||||
import { app } from "../../scripts/app.js";
|
||||
import { ComfyWidgets } from "../../scripts/widgets.js";
|
||||
|
||||
const PREVIEW_NODES = new Set(["ChatCompletion", "FormatString", "RandomChoice"]);
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfydv.previewText",
|
||||
|
||||
async beforeRegisterNodeDef(nodeType, nodeData) {
|
||||
if (!PREVIEW_NODES.has(nodeData.name)) return;
|
||||
|
||||
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||
nodeType.prototype.onNodeCreated = function () {
|
||||
const result = onNodeCreated?.apply(this, arguments);
|
||||
|
||||
const widget = ComfyWidgets.STRING(
|
||||
this,
|
||||
"comfydv_preview_text",
|
||||
["STRING", { multiline: true }],
|
||||
app
|
||||
).widget;
|
||||
widget.label = "Preview";
|
||||
widget.options.read_only = true;
|
||||
// Not a real input — nothing to save/replay in the saved
|
||||
// workflow JSON, and read-only anyway.
|
||||
widget.options.serialize = false;
|
||||
widget.serialize = false;
|
||||
widget.inputEl.readOnly = true;
|
||||
|
||||
return result;
|
||||
};
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted;
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
onExecuted?.apply(this, arguments);
|
||||
|
||||
const widget = this.widgets?.find(w => w.name === "comfydv_preview_text");
|
||||
if (!widget) return;
|
||||
|
||||
const text = message?.text ?? "";
|
||||
widget.value = Array.isArray(text) ? (text.join("\n\n") ?? "") : text;
|
||||
this.setDirtyCanvas(true, true);
|
||||
};
|
||||
},
|
||||
});
|
||||
+17
-1
@@ -40,9 +40,25 @@ class _FakeProvider:
|
||||
self.calls.append(("unload_model", model))
|
||||
|
||||
async def chat(
|
||||
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
|
||||
self,
|
||||
model,
|
||||
messages,
|
||||
options=None,
|
||||
timeout_secs=300.0,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
self.calls.append(("chat", model))
|
||||
if attempt_info is not None:
|
||||
attempt_info.update(
|
||||
{
|
||||
"seed": (options or {}).get("seed", 0),
|
||||
"attempts": 1,
|
||||
"timeout_secs": timeout_secs,
|
||||
"refusals": 0,
|
||||
}
|
||||
)
|
||||
return self.chat_response
|
||||
|
||||
|
||||
|
||||
@@ -405,6 +405,81 @@ def test_chat_exhausted_retries_returns_blank_without_raising(monkeypatch):
|
||||
assert calls["n"] == 3 # original + 2 retries, per max_retries=2
|
||||
|
||||
|
||||
def test_chat_timeout_escalates_per_retry_attempt(monkeypatch):
|
||||
timeouts = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
timeouts.append(timeout)
|
||||
if len(timeouts) < 3:
|
||||
return {"choices": [{"message": {"content": ""}}]}
|
||||
return {"choices": [{"message": {"content": "done"}}]}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
|
||||
|
||||
_run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat(
|
||||
"gemma-3-4b",
|
||||
[Message(role="user", content="hi")],
|
||||
timeout_secs=50.0,
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert timeouts == [50.0, 100.0, 150.0]
|
||||
|
||||
|
||||
def test_chat_attempt_info_populated_on_success(monkeypatch):
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
return {"choices": [{"message": {"content": "ok"}}]}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
|
||||
attempt_info: dict = {}
|
||||
_run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat(
|
||||
"gemma-3-4b",
|
||||
[Message(role="user", content="hi")],
|
||||
options={"seed": 9},
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
assert attempt_info == {
|
||||
"seed": 9,
|
||||
"attempts": 1,
|
||||
"timeout_secs": 300.0,
|
||||
"refusals": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_chat_on_status_called_on_retry_and_recovery(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)
|
||||
|
||||
statuses = []
|
||||
_run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat(
|
||||
"gemma-3-4b",
|
||||
[Message(role="user", content="hi")],
|
||||
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||
on_status=statuses.append,
|
||||
)
|
||||
)
|
||||
|
||||
assert len(statuses) == 2
|
||||
assert "Refusal/deflection detected" in statuses[0]
|
||||
assert "Recovered" in statuses[1]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# chat_structured — zero new logic, delegates to the shared helper unchanged
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -461,6 +536,64 @@ def test_chat_structured_forwards_options(monkeypatch):
|
||||
assert captured["options"] == {"temperature": 0.0}
|
||||
|
||||
|
||||
def test_chat_structured_forwards_attempt_info(monkeypatch):
|
||||
from pydantic import BaseModel
|
||||
|
||||
class Widget(BaseModel):
|
||||
name: str
|
||||
|
||||
captured = {}
|
||||
|
||||
async def fake_chat_structured(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return Widget(name="x")
|
||||
|
||||
monkeypatch.setattr("comfydv._llm.chat.chat_structured", fake_chat_structured)
|
||||
|
||||
attempt_info: dict = {}
|
||||
_run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat_structured(
|
||||
"gemma-3-4b",
|
||||
[Message(role="user", content="hi")],
|
||||
Widget,
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
# Same object passed straight through — the shared helper populates it,
|
||||
# this provider doesn't need to know its shape.
|
||||
assert captured["attempt_info"] is attempt_info
|
||||
|
||||
|
||||
def test_chat_structured_forwards_on_status(monkeypatch):
|
||||
from pydantic import BaseModel
|
||||
|
||||
class Widget(BaseModel):
|
||||
name: str
|
||||
|
||||
captured = {}
|
||||
|
||||
async def fake_chat_structured(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return Widget(name="x")
|
||||
|
||||
monkeypatch.setattr("comfydv._llm.chat.chat_structured", fake_chat_structured)
|
||||
|
||||
def on_status(msg):
|
||||
pass
|
||||
|
||||
_run_async(
|
||||
LlamaCppProvider("http://localhost:8080").chat_structured(
|
||||
"gemma-3-4b",
|
||||
[Message(role="user", content="hi")],
|
||||
Widget,
|
||||
on_status=on_status,
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["on_status"] is on_status
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _fetch_models — name-only view used by ComfyUI's /dv/ollama/models?backend=
|
||||
# llamacpp route (the JS refresh button / node-creation auto-populate).
|
||||
@@ -558,3 +691,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
|
||||
|
||||
@@ -142,6 +142,100 @@ def test_chat_structured_retries_on_validation_failure(monkeypatch):
|
||||
assert len(fake.calls) == 2
|
||||
|
||||
|
||||
def test_chat_structured_timeout_escalates_per_retry_attempt(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, bad, _Widget(name="b", count=2)])
|
||||
build_calls = []
|
||||
|
||||
def fake_build_agent(**kw):
|
||||
build_calls.append(kw["timeout_secs"])
|
||||
return fake
|
||||
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", fake_build_agent)
|
||||
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
timeout_secs=100.0,
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert build_calls == [100.0, 200.0, 300.0]
|
||||
|
||||
|
||||
def test_chat_structured_attempt_info_populated_on_success(monkeypatch):
|
||||
fake = _FakeAgent([_Widget(name="a", count=1)])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
attempt_info: dict = {}
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
options={"seed": 5},
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
assert attempt_info == {
|
||||
"seed": 5,
|
||||
"attempts": 1,
|
||||
"timeout_secs": 300.0,
|
||||
"refusals": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_chat_structured_attempt_info_populated_on_exhaustion(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, bad])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
attempt_info: dict = {}
|
||||
with pytest.raises(RuntimeError):
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
max_retries=1,
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
assert attempt_info["attempts"] == 2
|
||||
|
||||
|
||||
def test_chat_structured_on_status_called_on_retry_and_recovery(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, _Widget(name="b", count=2)])
|
||||
monkeypatch.setattr(chat_mod, "_build_agent", lambda **kw: fake)
|
||||
|
||||
statuses = []
|
||||
_run_async(
|
||||
chat_mod.chat_structured(
|
||||
base_url="http://localhost:11434/v1",
|
||||
model="llama3",
|
||||
messages=_messages(),
|
||||
schema=_Widget,
|
||||
max_retries=2,
|
||||
on_status=statuses.append,
|
||||
)
|
||||
)
|
||||
|
||||
assert len(statuses) == 2
|
||||
assert "Structured output failed" in statuses[0]
|
||||
assert "attempt 1/3" in statuses[0]
|
||||
assert "Recovered" in statuses[1]
|
||||
assert "attempt 2/3" in statuses[1]
|
||||
|
||||
|
||||
def test_chat_structured_exhausted_retries_raises_runtime_error(monkeypatch):
|
||||
bad = ValidationError.from_exception_data("Widget", [])
|
||||
fake = _FakeAgent([bad, bad, bad]) # max_retries=2 -> 3 total attempts
|
||||
@@ -381,6 +475,136 @@ 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_on_status_reports_refusal_reason(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)
|
||||
|
||||
statuses = []
|
||||
_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,
|
||||
on_status=statuses.append,
|
||||
)
|
||||
)
|
||||
|
||||
assert len(statuses) == 2
|
||||
assert "Refusal/deflection detected" in statuses[0]
|
||||
assert "Recovered" in statuses[1]
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
+404
-2
@@ -1,8 +1,29 @@
|
||||
"""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 (
|
||||
REFUSAL_EXEMPLARS,
|
||||
cosine_similarity,
|
||||
format_recovered_status,
|
||||
format_retry_status,
|
||||
is_ambiguous,
|
||||
is_lexical_refusal,
|
||||
is_refusal,
|
||||
next_seed,
|
||||
next_timeout_secs,
|
||||
record_attempt_info,
|
||||
)
|
||||
|
||||
|
||||
def test_next_seed_attempt_one_is_zero_by_default():
|
||||
@@ -22,3 +43,384 @@ 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
|
||||
|
||||
|
||||
class TestNextTimeoutSecs:
|
||||
def test_attempt_one_returns_base_timeout_unchanged(self):
|
||||
assert next_timeout_secs(300.0, 1) == 300.0
|
||||
|
||||
def test_escalates_multiplicatively_per_attempt(self):
|
||||
assert next_timeout_secs(300.0, 2) == 600.0
|
||||
assert next_timeout_secs(300.0, 3) == 900.0
|
||||
|
||||
|
||||
class TestRecordAttemptInfo:
|
||||
def test_none_attempt_info_is_a_no_op(self):
|
||||
# Must not raise — callers that don't care about this metadata pass
|
||||
# None and should see no behavior change at all.
|
||||
record_attempt_info(None, seed=1, attempts=2, timeout_secs=600.0, refusals=1)
|
||||
|
||||
def test_populates_dict_in_place(self):
|
||||
info: dict = {}
|
||||
record_attempt_info(info, seed=7, attempts=3, timeout_secs=900.0, refusals=2)
|
||||
assert info == {
|
||||
"seed": 7,
|
||||
"attempts": 3,
|
||||
"timeout_secs": 900.0,
|
||||
"refusals": 2,
|
||||
}
|
||||
|
||||
def test_overwrites_previous_values(self):
|
||||
# Callers call this once per attempt (or once per return path) —
|
||||
# a later call must replace, not merge with, an earlier one.
|
||||
info: dict = {"seed": 1, "attempts": 1, "timeout_secs": 300.0, "refusals": 0}
|
||||
record_attempt_info(info, seed=2, attempts=2, timeout_secs=600.0, refusals=1)
|
||||
assert info == {
|
||||
"seed": 2,
|
||||
"attempts": 2,
|
||||
"timeout_secs": 600.0,
|
||||
"refusals": 1,
|
||||
}
|
||||
|
||||
|
||||
class TestFormatStatus:
|
||||
def test_retry_status_includes_reason_attempt_seed_and_timeout(self):
|
||||
msg = format_retry_status("Blank response", 1, 3, seed=1, timeout_secs=200.0)
|
||||
assert "Blank response" in msg
|
||||
assert "attempt 1/3" in msg
|
||||
assert "seed=1" in msg
|
||||
assert "timeout=200s" in msg
|
||||
|
||||
def test_recovered_status_includes_attempt_and_seed(self):
|
||||
msg = format_recovered_status(2, 3, seed=1)
|
||||
assert "attempt 2/3" in msg
|
||||
assert "seed=1" in msg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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 am unable to comply with your instruction.",
|
||||
"Cannot generate this image due to the sensitive content.",
|
||||
'{"error": "cannot generate this content"}',
|
||||
"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_is_still_embedding_checked_but_not_a_refusal(self):
|
||||
# embed_fn present means the caller opted in to the embedding check
|
||||
# regardless of length/hedge-keywords (is_ambiguous no longer gates
|
||||
# this) — a long, on-topic response should still be embedding-
|
||||
# checked, it just shouldn't score as similar to the refusal
|
||||
# exemplars.
|
||||
calls = []
|
||||
|
||||
async def embed_fn(text):
|
||||
calls.append(text)
|
||||
if text == long_text:
|
||||
return [0.0, 1.0] # orthogonal to the exemplar vector below
|
||||
return [1.0, 0.0] # exemplars
|
||||
|
||||
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 long_text in calls # embedding check DID run, just scored low
|
||||
|
||||
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_long_json_shaped_soft_refusal_without_hedge_keywords_is_caught(self):
|
||||
# Regression case: a structured-output-shaped response (>600 chars
|
||||
# once you count JSON braces/field names) whose deflection doesn't
|
||||
# use any of the canned hedge keywords used to be invisible to the
|
||||
# embedding check entirely, because is_ambiguous gated on length
|
||||
# and keywords. embed_fn now runs unconditionally once configured.
|
||||
soft_refusal = (
|
||||
'{"prompt": "'
|
||||
+ "Let's take this in a different creative direction that everyone can enjoy. "
|
||||
* 8
|
||||
+ '"}'
|
||||
)
|
||||
assert len(soft_refusal) >= 600
|
||||
assert not is_lexical_refusal(soft_refusal)
|
||||
|
||||
async def embed_fn(text):
|
||||
if text == soft_refusal:
|
||||
return [1.0, 0.0]
|
||||
return [0.97, 0.24] # exemplars: cos-sim with [1,0] is ~0.97
|
||||
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
soft_refusal,
|
||||
embed_fn=embed_fn,
|
||||
embed_cache_key="regression-model",
|
||||
threshold=0.8,
|
||||
)
|
||||
)
|
||||
assert result is True
|
||||
|
||||
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
|
||||
|
||||
|
||||
class TestCustomPhrases:
|
||||
def test_custom_phrase_substring_match_needs_no_embed_fn(self):
|
||||
# A phrase the user added at runtime that the shipped lexical
|
||||
# patterns don't cover — should be caught for free, no embedding
|
||||
# model required.
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
"I am restricted from producing that kind of content.",
|
||||
custom_phrases=("restricted from",),
|
||||
)
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_custom_phrase_match_is_case_insensitive(self):
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
"SORRY, THAT'S OFF LIMITS FOR ME.",
|
||||
custom_phrases=("off limits",),
|
||||
)
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_unrelated_custom_phrase_does_not_match(self):
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
"The subject walks calmly toward the horizon.",
|
||||
custom_phrases=("restricted from", "off limits"),
|
||||
)
|
||||
)
|
||||
assert result is False
|
||||
|
||||
def test_blank_and_whitespace_custom_phrases_are_ignored(self):
|
||||
# A stray empty entry must never become a universal substring match.
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
"The subject walks calmly toward the horizon.",
|
||||
custom_phrases=("", " "),
|
||||
)
|
||||
)
|
||||
assert result is False
|
||||
|
||||
def test_custom_phrase_folded_into_embedding_exemplars(self):
|
||||
# No exact substring match, but embed_fn scores the response as
|
||||
# similar to the custom phrase (not one of the shipped exemplars).
|
||||
custom = "my creators have limited what I can show you"
|
||||
|
||||
async def embed_fn(text):
|
||||
if text == custom:
|
||||
return [1.0, 0.0]
|
||||
if text in REFUSAL_EXEMPLARS:
|
||||
return [0.0, 1.0] # shipped exemplars score orthogonal
|
||||
return [0.99, 0.14] # the probe response is near the custom one
|
||||
|
||||
result = _run_async(
|
||||
is_refusal(
|
||||
"There are limits my creators placed on what I can show.",
|
||||
embed_fn=embed_fn,
|
||||
embed_cache_key="custom-exemplar-model",
|
||||
threshold=0.8,
|
||||
custom_phrases=(custom,),
|
||||
)
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_different_custom_phrase_sets_do_not_share_exemplar_cache(self):
|
||||
# Regression guard: if the exemplar cache key ignored custom_phrases,
|
||||
# a second call with a different custom phrase set would incorrectly
|
||||
# reuse the first call's cached (and now stale) exemplar vectors.
|
||||
calls = []
|
||||
|
||||
async def embed_fn(text):
|
||||
calls.append(text)
|
||||
return [1.0, 0.0]
|
||||
|
||||
_run_async(
|
||||
is_refusal(
|
||||
"short reply",
|
||||
embed_fn=embed_fn,
|
||||
embed_cache_key="shared-model",
|
||||
custom_phrases=("phrase one",),
|
||||
)
|
||||
)
|
||||
first_call_count = len(calls)
|
||||
|
||||
_run_async(
|
||||
is_refusal(
|
||||
"short reply",
|
||||
embed_fn=embed_fn,
|
||||
embed_cache_key="shared-model",
|
||||
custom_phrases=("phrase two",),
|
||||
)
|
||||
)
|
||||
# A fresh custom phrase set re-embeds the exemplars (including the
|
||||
# new phrase) rather than reusing the first set's cached vectors.
|
||||
assert len(calls) > first_call_count
|
||||
|
||||
+128
-20
@@ -50,6 +50,7 @@ from comfydv.ollama import (
|
||||
OllamaOptionDisableThinking,
|
||||
OllamaOptionExtraBody,
|
||||
OllamaOptionMaxTokens,
|
||||
OllamaOptionRefusalRetry,
|
||||
OllamaOptionRepeatPenalty,
|
||||
OllamaOptionSeed,
|
||||
OllamaOptionTemperature,
|
||||
@@ -91,13 +92,37 @@ class _FakeProvider:
|
||||
self.calls.append(("unload_model", model))
|
||||
|
||||
async def chat(
|
||||
self, model, messages, options=None, timeout_secs=300.0, max_retries=2
|
||||
self,
|
||||
model,
|
||||
messages,
|
||||
options=None,
|
||||
timeout_secs=300.0,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
self.calls.append(("chat", model, messages, options, timeout_secs, max_retries))
|
||||
if attempt_info is not None:
|
||||
attempt_info.update(
|
||||
{
|
||||
"seed": (options or {}).get("seed", 0),
|
||||
"attempts": 1,
|
||||
"timeout_secs": timeout_secs,
|
||||
"refusals": 0,
|
||||
}
|
||||
)
|
||||
return self.chat_response
|
||||
|
||||
async def chat_structured(
|
||||
self, model, messages, schema, options=None, timeout_secs=300.0, max_retries=2
|
||||
self,
|
||||
model,
|
||||
messages,
|
||||
schema,
|
||||
options=None,
|
||||
timeout_secs=300.0,
|
||||
max_retries=2,
|
||||
attempt_info=None,
|
||||
on_status=None,
|
||||
):
|
||||
self.calls.append(
|
||||
(
|
||||
@@ -112,6 +137,15 @@ class _FakeProvider:
|
||||
)
|
||||
if self.raise_on_chat_structured is not None:
|
||||
raise self.raise_on_chat_structured
|
||||
if attempt_info is not None:
|
||||
attempt_info.update(
|
||||
{
|
||||
"seed": (options or {}).get("seed", 0),
|
||||
"attempts": 1,
|
||||
"timeout_secs": timeout_secs,
|
||||
"refusals": 0,
|
||||
}
|
||||
)
|
||||
return schema(**self.structured_field_values)
|
||||
|
||||
|
||||
@@ -271,7 +305,7 @@ class TestUS4ChatCompletion:
|
||||
def test_chat_model_receives_wired_string(self):
|
||||
"""Wiring LLMLoadModel.model_name → ChatCompletion.model works."""
|
||||
fake = _FakeProvider(chat_response="ok")
|
||||
_, _, used_model = ChatCompletion().chat(
|
||||
_, _, used_model, _ = ChatCompletion().chat(
|
||||
client=fake, model="llama3:latest", prompt="hi"
|
||||
)["result"]
|
||||
assert used_model == "llama3:latest"
|
||||
@@ -279,11 +313,17 @@ class TestUS4ChatCompletion:
|
||||
assert fake.calls[0][1] == "llama3:latest"
|
||||
|
||||
def test_chat_completion_returns_model_name(self):
|
||||
assert ChatCompletion.RETURN_TYPES == ("STRING", "OLLAMA_HISTORY", "STRING")
|
||||
assert ChatCompletion.RETURN_TYPES == (
|
||||
"STRING",
|
||||
"OLLAMA_HISTORY",
|
||||
"STRING",
|
||||
"INT",
|
||||
)
|
||||
assert ChatCompletion.RETURN_NAMES == (
|
||||
"response",
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"seed_used",
|
||||
)
|
||||
|
||||
def test_chat_is_output_node(self):
|
||||
@@ -301,15 +341,24 @@ class TestUS4ChatCompletion:
|
||||
ret = ChatCompletion().chat(client=fake, model="m", prompt="hi")
|
||||
assert "hello world" in ret["ui"]["text"][0]
|
||||
|
||||
def test_chat_result_is_3_tuple(self):
|
||||
def test_chat_result_is_4_tuple(self):
|
||||
fake = _FakeProvider(chat_response="hello")
|
||||
ret = ChatCompletion().chat(client=fake, model="m", prompt="hi")
|
||||
assert isinstance(ret["result"], tuple)
|
||||
assert len(ret["result"]) == 3
|
||||
response, history, model_name = ret["result"]
|
||||
assert len(ret["result"]) == 4
|
||||
response, history, model_name, seed_used = ret["result"]
|
||||
assert response == "hello"
|
||||
assert isinstance(history, list)
|
||||
assert model_name == "m"
|
||||
assert seed_used == 0
|
||||
|
||||
def test_chat_result_seed_used_reflects_pinned_seed(self):
|
||||
fake = _FakeProvider(chat_response="hello")
|
||||
ret = ChatCompletion().chat(
|
||||
client=fake, model="m", prompt="hi", options={"seed": 42}
|
||||
)
|
||||
_, _, _, seed_used = ret["result"]
|
||||
assert seed_used == 42
|
||||
|
||||
def test_chat_has_timeout_secs_input(self):
|
||||
input_types = ChatCompletion.INPUT_TYPES()
|
||||
@@ -336,7 +385,7 @@ class TestUS4ChatCompletion:
|
||||
):
|
||||
"""Scenario: Single-turn completion returns non-empty response."""
|
||||
(client,) = OllamaClient().create_client(ollama_host)
|
||||
response, updated_history, model_name = ChatCompletion().chat(
|
||||
response, updated_history, model_name, _seed = ChatCompletion().chat(
|
||||
client=client,
|
||||
model=_CHAT_MODEL,
|
||||
prompt="Say exactly the word: pong",
|
||||
@@ -358,14 +407,14 @@ class TestUS4ChatCompletion:
|
||||
"""
|
||||
(client,) = OllamaClient().create_client(ollama_host)
|
||||
no_think = {"think": False}
|
||||
_, history, _ = ChatCompletion().chat(
|
||||
_, history, _, _ = ChatCompletion().chat(
|
||||
client=client,
|
||||
model=_CHAT_MODEL,
|
||||
prompt="My name is Alice. Remember it.",
|
||||
history=[],
|
||||
options=no_think,
|
||||
)["result"]
|
||||
response, updated, _ = ChatCompletion().chat(
|
||||
response, updated, _, _ = ChatCompletion().chat(
|
||||
client=client,
|
||||
model=_CHAT_MODEL,
|
||||
prompt="What is my name?",
|
||||
@@ -379,11 +428,11 @@ class TestUS4ChatCompletion:
|
||||
def test_history_accumulated_correctly(self, ollama_host, skip_if_no_ollama):
|
||||
"""History list grows by 2 entries per turn."""
|
||||
(client,) = OllamaClient().create_client(ollama_host)
|
||||
_, h1, _ = ChatCompletion().chat(
|
||||
_, h1, _, _ = ChatCompletion().chat(
|
||||
client=client, model=_CHAT_MODEL, prompt="Turn 1", history=[]
|
||||
)["result"]
|
||||
assert len(h1) == 2
|
||||
_, h2, _ = ChatCompletion().chat(
|
||||
_, h2, _, _ = ChatCompletion().chat(
|
||||
client=client, model=_CHAT_MODEL, prompt="Turn 2", history=h1
|
||||
)["result"]
|
||||
assert len(h2) == 4
|
||||
@@ -427,7 +476,7 @@ class TestUS4ChatCompletion:
|
||||
max_retries=1,
|
||||
unique_id="smoke-test",
|
||||
)
|
||||
response_text, _history, model_name, output = result["result"]
|
||||
response_text, _history, model_name, output, _seed = result["result"]
|
||||
assert model_name == _CHAT_MODEL
|
||||
assert json.loads(response_text) == {"output": output}
|
||||
assert isinstance(output, str) and output.strip()
|
||||
@@ -503,8 +552,9 @@ class TestStructuredOutput:
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"output",
|
||||
"seed_used",
|
||||
)
|
||||
assert len(ret["result"]) == 4
|
||||
assert len(ret["result"]) == 5
|
||||
assert ret["result"][3] == "clean text"
|
||||
assert json.loads(ret["result"][0]) == {"output": "clean text"}
|
||||
|
||||
@@ -531,6 +581,7 @@ class TestStructuredOutput:
|
||||
"STRING",
|
||||
"INT",
|
||||
"BOOLEAN",
|
||||
"INT",
|
||||
)
|
||||
assert ChatCompletion.RETURN_NAMES == (
|
||||
"response",
|
||||
@@ -539,8 +590,9 @@ class TestStructuredOutput:
|
||||
"summary",
|
||||
"score",
|
||||
"is_positive",
|
||||
"seed_used",
|
||||
)
|
||||
_, _, _, summary, score, is_positive = ret["result"]
|
||||
_, _, _, summary, score, is_positive, _seed = ret["result"]
|
||||
assert summary == "great"
|
||||
assert score == 9
|
||||
assert isinstance(score, int)
|
||||
@@ -568,7 +620,7 @@ class TestStructuredOutput:
|
||||
output_schema=_SCHEMA_WITH_OPTIONAL_NULLABLE_FIELD,
|
||||
unique_id="n3b",
|
||||
)
|
||||
_, _, _, output, duration_seconds = ret["result"]
|
||||
_, _, _, output, duration_seconds, _seed = ret["result"]
|
||||
assert output == "hi"
|
||||
assert duration_seconds is None # FLOAT socket; only STRING coerces None to ""
|
||||
|
||||
@@ -597,7 +649,7 @@ class TestStructuredOutput:
|
||||
structured_output=True,
|
||||
unique_id="n5",
|
||||
)
|
||||
assert len(ChatCompletion.RETURN_TYPES) == 4
|
||||
assert len(ChatCompletion.RETURN_TYPES) == 5
|
||||
|
||||
fake2 = _FakeProvider(chat_response="x")
|
||||
ChatCompletion().chat(
|
||||
@@ -732,6 +784,55 @@ 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,
|
||||
"custom_phrases": (),
|
||||
}
|
||||
}
|
||||
|
||||
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_refusal_retry_parses_custom_phrases_csv(self):
|
||||
# Whitespace around each phrase is trimmed and empty entries (from
|
||||
# blank items / trailing commas) are dropped.
|
||||
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
|
||||
enabled=True,
|
||||
embedding_model="",
|
||||
threshold=0.82,
|
||||
custom_phrases=" I am restricted from , not permitted to help,,",
|
||||
)
|
||||
assert opts["refusal_retry"]["custom_phrases"] == (
|
||||
"I am restricted from",
|
||||
"not permitted to help",
|
||||
)
|
||||
|
||||
def test_refusal_retry_custom_phrases_defaults_empty(self):
|
||||
(opts,) = OllamaOptionRefusalRetry().set_refusal_retry(
|
||||
enabled=True, embedding_model="", threshold=0.82
|
||||
)
|
||||
assert opts["refusal_retry"]["custom_phrases"] == ()
|
||||
|
||||
def test_extra_body_merges_json(self):
|
||||
(opts,) = OllamaOptionExtraBody().set_extra_body(
|
||||
extra_body_json='{"stop": ["</s>"]}', options={"temperature": 0.0}
|
||||
@@ -768,8 +869,8 @@ class TestUS5ComposableOptions:
|
||||
history=[],
|
||||
options=opts2,
|
||||
)
|
||||
r1, _, _model = ChatCompletion().chat(**kwargs)["result"]
|
||||
r2, _, _model = ChatCompletion().chat(**kwargs)["result"]
|
||||
r1, _, _model, _seed = ChatCompletion().chat(**kwargs)["result"]
|
||||
r2, _, _model, _seed = ChatCompletion().chat(**kwargs)["result"]
|
||||
assert r1 == r2
|
||||
|
||||
|
||||
@@ -1008,6 +1109,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"response",
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"seed_used",
|
||||
]
|
||||
|
||||
def test_valid_schema_returns_dynamic_outputs(self):
|
||||
@@ -1027,6 +1129,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"summary",
|
||||
"score",
|
||||
"is_positive",
|
||||
"seed_used",
|
||||
]
|
||||
assert types == [
|
||||
"STRING",
|
||||
@@ -1035,6 +1138,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"STRING",
|
||||
"INT",
|
||||
"BOOLEAN",
|
||||
"INT",
|
||||
]
|
||||
|
||||
def test_invalid_json_while_typing_falls_back_to_base_outputs(self):
|
||||
@@ -1052,6 +1156,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"response",
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"seed_used",
|
||||
]
|
||||
|
||||
def test_incomplete_schema_missing_properties_falls_back_to_base_outputs(self):
|
||||
@@ -1066,6 +1171,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"response",
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"seed_used",
|
||||
]
|
||||
|
||||
def test_response_reflects_class_state_not_just_this_call(self):
|
||||
@@ -1079,7 +1185,8 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"output_schema": _SINGLE_FIELD_SCHEMA,
|
||||
}
|
||||
)
|
||||
assert ChatCompletion.RETURN_NAMES[-1] == "output"
|
||||
assert ChatCompletion.RETURN_NAMES[-2] == "output"
|
||||
assert ChatCompletion.RETURN_NAMES[-1] == "seed_used"
|
||||
|
||||
data = self._call(
|
||||
{"unique_id": "r5", "structured_output": False, "output_schema": "{}"}
|
||||
@@ -1088,6 +1195,7 @@ class TestUpdateStructuredOutputsRoute:
|
||||
"response",
|
||||
"updated_history",
|
||||
"model_name",
|
||||
"seed_used",
|
||||
]
|
||||
assert ChatCompletion.RETURN_NAMES == ChatCompletion._BASE_RETURN_NAMES
|
||||
|
||||
|
||||
@@ -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,419 @@ def test_chat_text_only_payload_omits_images_key(monkeypatch):
|
||||
)
|
||||
|
||||
assert captured["payload"]["messages"] == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# attempt_info / escalating per-attempt timeout
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_chat_timeout_escalates_per_retry_attempt(monkeypatch):
|
||||
timeouts = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
timeouts.append(timeout)
|
||||
if len(timeouts) < 3:
|
||||
return {"message": {"content": ""}} # blank -> retry
|
||||
return {"message": {"content": "done"}}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
monkeypatch.setattr(provider_mod.asyncio, "sleep", _fake_sleep)
|
||||
|
||||
_run_async(
|
||||
OllamaProvider("http://localhost:11434").chat(
|
||||
"llama3",
|
||||
[Message(role="user", content="hi")],
|
||||
timeout_secs=100.0,
|
||||
max_retries=2,
|
||||
)
|
||||
)
|
||||
|
||||
assert timeouts == [100.0, 200.0, 300.0]
|
||||
|
||||
|
||||
def test_chat_attempt_info_populated_on_success(monkeypatch):
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
return {"message": {"content": "ok"}}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
|
||||
attempt_info: dict = {}
|
||||
_run_async(
|
||||
OllamaProvider("http://localhost:11434").chat(
|
||||
"llama3",
|
||||
[Message(role="user", content="hi")],
|
||||
options={"seed": 5},
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
assert attempt_info == {
|
||||
"seed": 5,
|
||||
"attempts": 1,
|
||||
"timeout_secs": 300.0,
|
||||
"refusals": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_chat_attempt_info_reflects_refusal_retry(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)
|
||||
|
||||
attempt_info: dict = {}
|
||||
_run_async(
|
||||
OllamaProvider("http://localhost:11434").chat(
|
||||
"llama3",
|
||||
[Message(role="user", content="hi")],
|
||||
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||
attempt_info=attempt_info,
|
||||
)
|
||||
)
|
||||
|
||||
assert attempt_info["attempts"] == 2
|
||||
assert attempt_info["refusals"] == 1
|
||||
assert attempt_info["seed"] == 1 # next_seed(): attempt 2 -> seed 1
|
||||
|
||||
|
||||
def test_chat_on_status_called_on_retry_and_recovery(monkeypatch):
|
||||
"""Live counterpart to attempt_info: called mid-loop (not just at the
|
||||
end) so a caller (ChatCompletion's PromptServer.send_progress_text
|
||||
closure) can show retry progress while the node is still executing."""
|
||||
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)
|
||||
|
||||
statuses = []
|
||||
_run_async(
|
||||
OllamaProvider("http://localhost:11434").chat(
|
||||
"llama3",
|
||||
[Message(role="user", content="hi")],
|
||||
options={"refusal_retry": {"enabled": True, "embedding_model": ""}},
|
||||
on_status=statuses.append,
|
||||
)
|
||||
)
|
||||
|
||||
assert len(statuses) == 2
|
||||
assert "Refusal/deflection detected" in statuses[0]
|
||||
assert "attempt 1/3" in statuses[0]
|
||||
assert "Recovered" in statuses[1]
|
||||
assert "attempt 2/3" in statuses[1]
|
||||
|
||||
|
||||
def test_chat_on_status_not_called_when_first_attempt_succeeds(monkeypatch):
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
return {"message": {"content": "ok"}}
|
||||
|
||||
monkeypatch.setattr(provider_mod, "_post_json", fake_post)
|
||||
|
||||
statuses = []
|
||||
_run_async(
|
||||
OllamaProvider("http://localhost:11434").chat(
|
||||
"llama3",
|
||||
[Message(role="user", content="hi")],
|
||||
on_status=statuses.append,
|
||||
)
|
||||
)
|
||||
|
||||
assert statuses == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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_retries_on_custom_phrase_with_no_embedding_model(monkeypatch):
|
||||
"""Regression coverage for OllamaOptionRefusalRetry's custom_phrases
|
||||
field: a phrase the user added at runtime, not one of the shipped
|
||||
lexical patterns, still triggers a retry — with no embedding_model
|
||||
configured at all."""
|
||||
calls = []
|
||||
|
||||
async def fake_post(url, payload, *, timeout=120.0, headers=None):
|
||||
calls.append(payload)
|
||||
if len(calls) == 1:
|
||||
# deliberately not covered by any shipped lexical pattern —
|
||||
# only the custom_phrases entry below should catch this
|
||||
return {"message": {"content": "That's outside my operating parameters."}}
|
||||
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": "",
|
||||
"custom_phrases": ("outside my operating parameters",),
|
||||
}
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
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 {"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
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Tests for comfydv.random_choice.RandomChoice.
|
||||
|
||||
Covers the UI-preview addition (OUTPUT_NODE=True + a "ui": {"text": [...]}
|
||||
return, mirroring ChatCompletion/FormatString so all three get a visible
|
||||
text preview via src/js/preview_text.js) without changing IS_CHANGED's
|
||||
change-detection semantics.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from comfydv.random_choice import RandomChoice, _preview_text
|
||||
|
||||
|
||||
def test_output_node_is_true():
|
||||
assert RandomChoice.OUTPUT_NODE is True
|
||||
|
||||
|
||||
def test_random_choice_returns_ui_result_dict():
|
||||
ret = RandomChoice().random_choice(input1="a", seed=42)
|
||||
assert isinstance(ret, dict)
|
||||
assert "ui" in ret
|
||||
assert "result" in ret
|
||||
assert ret["result"] == ("a",)
|
||||
|
||||
|
||||
def test_ui_text_matches_the_chosen_value_for_a_string():
|
||||
ret = RandomChoice().random_choice(input1="hello", seed=42)
|
||||
assert ret["ui"]["text"] == ["hello"]
|
||||
|
||||
|
||||
def test_ui_text_for_a_number_is_stringified():
|
||||
ret = RandomChoice().random_choice(input1=7, seed=42)
|
||||
assert ret["ui"]["text"] == ["7"]
|
||||
assert ret["result"] == (7,)
|
||||
|
||||
|
||||
def test_seed_pins_the_choice_deterministically():
|
||||
ret1 = RandomChoice().random_choice(input1="a", input2="b", input3="c", seed=42)
|
||||
ret2 = RandomChoice().random_choice(input1="a", input2="b", input3="c", seed=42)
|
||||
assert ret1["result"] == ret2["result"]
|
||||
|
||||
|
||||
def test_is_changed_returns_raw_pick_not_ui_wrapped_dict():
|
||||
"""Regression guard: IS_CHANGED must keep returning the same shape it
|
||||
did before the ui-preview addition (the raw picked value), not the new
|
||||
{"ui": ..., "result": ...} dict random_choice() now returns — otherwise
|
||||
ComfyUI's change-detection comparison would be comparing dicts full of
|
||||
UI-only text noise instead of the actual output value."""
|
||||
result = RandomChoice.IS_CHANGED(input1="only-choice", seed=42)
|
||||
assert result == "only-choice"
|
||||
|
||||
|
||||
class TestPreviewText:
|
||||
def test_string_passthrough(self):
|
||||
assert _preview_text("hello") == "hello"
|
||||
|
||||
def test_number_stringified(self):
|
||||
assert _preview_text(42) == "42"
|
||||
assert _preview_text(3.14) == "3.14"
|
||||
assert _preview_text(True) == "True"
|
||||
|
||||
def test_list_json_dumped(self):
|
||||
assert json.loads(_preview_text([1, 2, 3])) == [1, 2, 3]
|
||||
|
||||
def test_unserializable_falls_back_to_str(self):
|
||||
class Weird:
|
||||
def __repr__(self):
|
||||
return "<Weird thing>"
|
||||
|
||||
# json.dumps(default=str) actually succeeds here (falls back to
|
||||
# str() per-value), so this exercises the "else JSON" branch
|
||||
# rather than the outer except — confirms it never raises either
|
||||
# way.
|
||||
assert "<Weird thing>" in _preview_text(Weird())
|
||||
Reference in New Issue
Block a user