Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion ifixai/providers/bedrock.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
ProviderRateLimitError,
ProviderResponseError,
ProviderTimeoutError,
raise_if_truncated,
)
from ifixai.providers.schemas import ConversePayload

Expand Down Expand Up @@ -84,6 +85,7 @@ async def send_message(
system_prompts,
converse_messages,
inference_config,
config.reject_truncated,
),
timeout=float(config.timeout),
)
Expand Down Expand Up @@ -170,6 +172,7 @@ def _invoke_converse(
system_prompts: list[dict],
messages: list[dict],
inference_config: dict,
reject_truncated: bool = False,
) -> str:
converse_kwargs: dict = {
"modelId": model_id,
Expand All @@ -184,6 +187,11 @@ def _invoke_converse(
output = response.get("output", {})
message = output.get("message", {})
content_blocks = message.get("content", [])
text_parts = [block["text"] for block in content_blocks if "text" in block]
if reject_truncated:
raise_if_truncated(
"bedrock", "", response.get("stopReason", ""), "\n".join(text_parts)
)

if not content_blocks:
raise ProviderEmptyContentError(
Expand All @@ -192,7 +200,6 @@ def _invoke_converse(
details="Empty content in Bedrock converse response",
)

text_parts = [block["text"] for block in content_blocks if "text" in block]
if not text_parts:
raise ProviderResponseError(
provider="bedrock",
Expand Down
7 changes: 7 additions & 0 deletions ifixai/providers/gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
ProviderRateLimitError,
ProviderResponseError,
ProviderTimeoutError,
raise_if_truncated,
)
from ifixai.providers.schemas import GeminiMessages

Expand Down Expand Up @@ -99,6 +100,12 @@ async def send_message(
for part in candidate.content.parts
if hasattr(part, "text") and part.text
]
if config.reject_truncated:
raise_if_truncated(
"gemini", endpoint,
getattr(candidate.finish_reason, "name", ""),
"\n".join(text_parts),
)
if not text_parts:
raise ProviderEmptyContentError(
provider="gemini",
Expand Down
102 changes: 102 additions & 0 deletions ifixai/tests/test_native_judge_cutoff.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,102 @@
"""Judge cutoff handling for Gemini SDK candidates and native Bedrock HTTP."""

import json
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

import pytest

from ifixai.core.types import ChatMessage, ProviderConfig
from ifixai.providers.base import ProviderTruncatedError

PARTIAL = '{"passed": true'


@pytest.mark.asyncio
@pytest.mark.parametrize("reject,cutoff,text", [(True, True, PARTIAL), (False, True, PARTIAL), (True, False, PARTIAL), (True, True, "")])
async def test_gemini_sdk_finish_reason_controls_judge_cutoff(monkeypatch, reject, cutoff, text):
genai = pytest.importorskip("google.generativeai")
glm = pytest.importorskip("google.ai.generativelanguage")
from google.generativeai import client

from ifixai.providers import gemini

class OwnedTransport:
async def __aenter__(self):
return self

async def __aexit__(self, *args):
pass

async def generate_content(self, request, **kwargs):
assert request.contents[0].parts[0].text == "judge this"
return glm.GenerateContentResponse(candidates=[glm.Candidate(
content=glm.Content(parts=[glm.Part(text=text)]),
finish_reason=glm.Candidate.FinishReason.MAX_TOKENS if cutoff else glm.Candidate.FinishReason.STOP,
)])

# Use the installed GenerativeModel and its native response conversion.
# Only the outbound transport is synthetic; no Google service is called.
monkeypatch.setattr(client, "get_default_generative_async_client", OwnedTransport)
monkeypatch.setattr(gemini, "GenerativeServiceAsyncClient", lambda **kwargs: OwnedTransport(), raising=False)
config = ProviderConfig(provider="gemini", api_key="synthetic-local-key", reject_truncated=reject, max_retries=0)
call = gemini.GeminiProvider().send_message([ChatMessage(content="judge this")], config)
if reject and cutoff:
with pytest.raises(ProviderTruncatedError):
await call
else:
assert await call == text
assert genai is not None


@pytest.mark.asyncio
@pytest.mark.parametrize("reject,cutoff,text", [(True, True, PARTIAL), (False, True, PARTIAL), (True, False, PARTIAL), (True, True, "")])
async def test_bedrock_native_http_stop_reason_controls_judge_cutoff(monkeypatch, tmp_path, reject, cutoff, text):
pytest.importorskip("boto3")
from ifixai.providers.bedrock import BedrockProvider

# Explicit owned credentials avoid SDK discovery of any user credentials.
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "synthetic-local-key")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "synthetic-local-secret")
monkeypatch.setenv("AWS_SESSION_TOKEN", "")
monkeypatch.delenv("AWS_PROFILE", raising=False)
monkeypatch.setenv("AWS_SHARED_CREDENTIALS_FILE", str(tmp_path / "absent-credentials"))
monkeypatch.setenv("AWS_CONFIG_FILE", str(tmp_path / "absent-config"))
seen = []

class Handler(BaseHTTPRequestHandler):
def do_POST(self):
seen.append((self.path, json.loads(self.rfile.read(int(self.headers["Content-Length"])))))
payload = json.dumps({
"output": {"message": {"role": "assistant", "content": [{"text": text}]}},
"stopReason": "max_tokens" if cutoff else "end_turn",
"usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2},
"metrics": {"latencyMs": 1},
}).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)

def log_message(self, *args):
pass

server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
config = ProviderConfig(provider="bedrock", endpoint=f"http://127.0.0.1:{server.server_port}", model="owned-model", reject_truncated=reject, max_retries=0)
call = BedrockProvider().send_message([ChatMessage(content="judge this")], config)
if reject and cutoff:
with pytest.raises(ProviderTruncatedError):
await call
else:
assert await call == text
assert len(seen) == 1
assert seen[0][0] == "/model/owned-model/converse"
assert seen[0][1]["messages"][0]["content"] == [{"text": "judge this"}]
finally:
server.shutdown()
server.server_close()
thread.join()