diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 8b7ad502..fc9a1f76 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -1,4 +1,5 @@ import datetime +import json import logging import time from typing import Any @@ -43,6 +44,52 @@ logger = logging.getLogger(__name__) logging.getLogger("sqlalchemy.engine.Engine").disabled = True +def repair_json(json_str): + """Attempt to repair incomplete JSON by adding missing closing braces/brackets""" + json_str = json_str.strip() + + # Count opening vs closing braces and brackets + open_braces = json_str.count("{") + close_braces = json_str.count("}") + open_brackets = json_str.count("[") + close_brackets = json_str.count("]") + + # Add missing closing characters + missing_brackets = open_brackets - close_brackets + missing_braces = open_braces - close_braces + + repaired = json_str + repaired += "]" * missing_brackets + repaired += "}" * missing_braces + + return repaired + + +def validate_and_repair_json(json_str): + """Validate JSON and attempt repairs if needed""" + json_str = json_str.strip() + + try: + # Try parsing as-is first + json.loads(json_str) + return json_str + except json.JSONDecodeError as e: + print(f"JSON error: {e}") + print(f"Error at position: {e.pos}") + print(f"Context: ...{json_str[max(0, e.pos - 50) : e.pos + 50]}...") + + # Attempt repair + repaired = repair_json(json_str) + + try: + json.loads(repaired) + return repaired + except json.JSONDecodeError: + raise ValueError( + f"Could not repair JSON. Original length: {len(json_str)}, Repaired length: {len(repaired)}" + ) + + @honcho_llm_call( provider=settings.DERIVER.PROVIDER, model=settings.DERIVER.MODEL, @@ -55,6 +102,15 @@ logging.getLogger("sqlalchemy.engine.Engine").disabled = True else None, enable_retry=True, retry_attempts=3, + response_format={ + "type": "json_schema", + "json_schema": { + "name": ReasoningResponse.__name__, + "schema": ReasoningResponse.model_json_schema(), + }, + }, + # if settings.DERIVER.PROVIDER == "custom" + # else None, # Only for vllm/custom provider ) async def critical_analysis_call( peer_name: str, @@ -364,8 +420,16 @@ class CertaintyReasoner: formatted_new_turn, ) - # Call the standalone LLM function (now with Tenacity retries) - response_obj = await critical_analysis_call( + # Raw Debugging logic + + import openai + + client = openai.OpenAI( + base_url=settings.LLM.OPENAI_COMPATIBLE_BASE_URL, + api_key=settings.LLM.OPENAI_COMPATIBLE_API_KEY, + ) + + prompt = critical_analysis_prompt( peer_name=speaker, message_created_at=message_created_at, context=formatted_context, @@ -373,6 +437,54 @@ class CertaintyReasoner: new_turn=formatted_new_turn, ) + logger.debug("CRITICAL ANALYSIS: Start") + + response = client.chat.completions.create( + model=settings.DERIVER.MODEL, + messages=prompt, + stop=[" \n", "\n\n\n\n"], + max_tokens=settings.DERIVER.MAX_OUTPUT_TOKENS, + response_format={ + "type": "json_schema", + "json_schema": { + "name": ReasoningResponse.__name__, + "schema": ReasoningResponse.model_json_schema(), + }, + }, + ) + + print("============================= test_rep ===============================") + + test_rep = response.choices[0].message.content + print(test_rep) + + logger.debug("CRITICAL ANALYSIS: Finished") + print( + "============================= test_rep finished ===============================" + ) + + final = validate_and_repair_json(test_rep) + + response_obj = ReasoningResponse.model_validate_json(final) + + # Call the standalone LLM function (now with Tenacity retries) + # response_obj = await critical_analysis_call( + # peer_name=speaker, + # message_created_at=message_created_at, + # context=formatted_context, + # history=history, + # new_turn=formatted_new_turn, + # ) + + # print("================= Response Object =================") + # print(type(response_obj)) + # print(response_obj.__dict__) + # print(list(response_obj.model_fields.keys())) + # print(response_obj.response.content) + # print(response_obj.content) + # ['metadata', 'response', 'tool_types', 'prompt_template', 'fn_args', 'dynamic_config', 'messages', 'call_params', 'call_kwargs', 'user_message_param', 'start_time', 'end_time'] + # print("================= Response Object End =================") + # Handle different response types if isinstance(response_obj, str): # If response is a string, try to parse as JSON diff --git a/src/utils/clients.py b/src/utils/clients.py index 7d0ea357..5e60b0fb 100644 --- a/src/utils/clients.py +++ b/src/utils/clients.py @@ -207,6 +207,7 @@ def honcho_llm_call( retry_attempts: int = 3, stream: bool = False, return_call_response: bool = False, # pyright: ignore + response_format: dict[str, Any] | None = None, # Add response_format parameter **extra_call_params: Any, ) -> Any: """ @@ -322,6 +323,7 @@ def honcho_llm_call( # https://mirascope.com/docs/mirascope/learn/provider-specific/openai#response-models if resolved_provider == "openai": response_model.model_config = ResponseModelConfigDict(strict=True) + llm_kwargs["response_model"] = response_model if json_mode: llm_kwargs["json_mode"] = json_mode