fix: checkpoint
This commit is contained in:
parent
550c33def8
commit
b7c1f5c94c
|
|
@ -181,3 +181,6 @@ timing_logs.csv
|
|||
|
||||
config.toml
|
||||
.aider*
|
||||
|
||||
.crush/
|
||||
CRUSH.md
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from src import crud
|
|||
from src.config import settings
|
||||
from src.dependencies import tracked_db
|
||||
from src.utils import summarizer
|
||||
from src.utils.clients import honcho_llm_call
|
||||
from src.utils.clients import create_retry_wrapper, direct_llm_call
|
||||
from src.utils.embedding_store import EmbeddingStore
|
||||
from src.utils.formatting import (
|
||||
REASONING_LEVELS,
|
||||
|
|
@ -97,36 +97,15 @@ def validate_and_repair_json(json_str: str):
|
|||
) from e
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.DERIVER.PROVIDER,
|
||||
model=settings.DERIVER.MODEL,
|
||||
track_name="Critical Analysis Call",
|
||||
response_model=ReasoningResponse,
|
||||
json_mode=True,
|
||||
max_tokens=settings.DERIVER.MAX_OUTPUT_TOKENS or settings.LLM.DEFAULT_MAX_TOKENS,
|
||||
thinking_budget_tokens=settings.DERIVER.THINKING_BUDGET_TOKENS
|
||||
if settings.DERIVER.PROVIDER == "anthropic"
|
||||
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
|
||||
)
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def critical_analysis_call(
|
||||
peer_name: str,
|
||||
message_created_at: datetime.datetime,
|
||||
context: str,
|
||||
history: str,
|
||||
new_turn: str,
|
||||
):
|
||||
return critical_analysis_prompt(
|
||||
) -> ReasoningResponse:
|
||||
prompt_content = critical_analysis_prompt(
|
||||
peer_name=peer_name,
|
||||
message_created_at=message_created_at,
|
||||
context=context,
|
||||
|
|
@ -134,6 +113,22 @@ async def critical_analysis_call(
|
|||
new_turn=new_turn,
|
||||
)
|
||||
|
||||
response = await direct_llm_call(
|
||||
prompt=prompt_content,
|
||||
provider=settings.DERIVER.PROVIDER,
|
||||
model=settings.DERIVER.MODEL,
|
||||
response_model=ReasoningResponse,
|
||||
json_mode=True,
|
||||
max_tokens=settings.DERIVER.MAX_OUTPUT_TOKENS
|
||||
or settings.LLM.DEFAULT_MAX_TOKENS,
|
||||
thinking_budget_tokens=settings.DERIVER.THINKING_BUDGET_TOKENS
|
||||
if settings.DERIVER.PROVIDER == "anthropic"
|
||||
else None,
|
||||
track_name="Critical Analysis Call",
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@conditional_observe
|
||||
class Deriver:
|
||||
|
|
|
|||
|
|
@ -9,17 +9,15 @@ historical observations.
|
|||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import tiktoken
|
||||
from dotenv import load_dotenv
|
||||
from langfuse.decorators import langfuse_context
|
||||
from mirascope.llm import Stream
|
||||
|
||||
from src import crud
|
||||
from src.config import settings
|
||||
from src.dependencies import tracked_db
|
||||
from src.routers.sessions import get_session_context
|
||||
from src.utils.clients import honcho_llm_call
|
||||
from src.utils.clients import create_retry_wrapper, direct_llm_call
|
||||
from src.utils.embedding_store import EmbeddingStore
|
||||
from src.utils.logging import (
|
||||
accumulate_metric,
|
||||
|
|
@ -32,21 +30,13 @@ from .utils import get_observations
|
|||
# Configure logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Load environment variables
|
||||
load_dotenv()
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.DIALECTIC.PROVIDER,
|
||||
model=settings.DIALECTIC.MODEL,
|
||||
track_name="Dialectic Call",
|
||||
max_tokens=settings.DIALECTIC.MAX_OUTPUT_TOKENS,
|
||||
thinking_budget_tokens=settings.DIALECTIC.THINKING_BUDGET_TOKENS
|
||||
if settings.DIALECTIC.PROVIDER == "anthropic"
|
||||
else None,
|
||||
enable_retry=True,
|
||||
retry_attempts=3,
|
||||
)
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def dialectic_call(
|
||||
query: str,
|
||||
working_representation: str | None,
|
||||
|
|
@ -67,42 +57,36 @@ async def dialectic_call(
|
|||
Returns:
|
||||
Model response
|
||||
"""
|
||||
# Generate the prompt and log it
|
||||
prompt_result = dialectic_prompt(
|
||||
query,
|
||||
working_representation,
|
||||
recent_conversation_history,
|
||||
additional_context,
|
||||
peer_name,
|
||||
target_name,
|
||||
# Generate the prompt
|
||||
prompt_content = dialectic_prompt(
|
||||
query=query,
|
||||
working_representation=working_representation,
|
||||
recent_conversation_history=recent_conversation_history,
|
||||
additional_context=additional_context,
|
||||
peer_name=peer_name,
|
||||
target_name=target_name,
|
||||
)
|
||||
|
||||
# Pretty print the prompt content
|
||||
if len(prompt_result) > 0:
|
||||
# Extract content from the first BaseMessageParam
|
||||
prompt_content = prompt_result[0].content
|
||||
else:
|
||||
prompt_content = str(prompt_result)
|
||||
|
||||
logger.debug("=== DIALECTIC PROMPT ===")
|
||||
logger.debug(prompt_content)
|
||||
logger.debug("=== END DIALECTIC PROMPT ===")
|
||||
|
||||
return prompt_result
|
||||
# Make direct LLM call
|
||||
response = await direct_llm_call(
|
||||
prompt=prompt_content,
|
||||
provider=settings.DIALECTIC.PROVIDER,
|
||||
model=settings.DIALECTIC.MODEL,
|
||||
max_tokens=settings.DIALECTIC.MAX_OUTPUT_TOKENS,
|
||||
thinking_budget_tokens=settings.DIALECTIC.THINKING_BUDGET_TOKENS
|
||||
if settings.DIALECTIC.PROVIDER == "anthropic"
|
||||
else None,
|
||||
track_name="Dialectic Call",
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.DIALECTIC.PROVIDER,
|
||||
model=settings.DIALECTIC.MODEL,
|
||||
track_name="Dialectic Stream",
|
||||
max_tokens=settings.DIALECTIC.MAX_OUTPUT_TOKENS,
|
||||
thinking_budget_tokens=settings.DIALECTIC.THINKING_BUDGET_TOKENS
|
||||
if settings.DIALECTIC.PROVIDER == "anthropic"
|
||||
else None,
|
||||
enable_retry=True,
|
||||
retry_attempts=3,
|
||||
stream=True,
|
||||
)
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def dialectic_stream(
|
||||
query: str,
|
||||
working_representation: str | None,
|
||||
|
|
@ -110,7 +94,7 @@ async def dialectic_stream(
|
|||
additional_context: str | None,
|
||||
peer_name: str,
|
||||
target_name: str | None = None,
|
||||
):
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Make a streaming call to the dialectic model for context synthesis.
|
||||
|
||||
|
|
@ -123,28 +107,35 @@ async def dialectic_stream(
|
|||
Returns:
|
||||
Streaming model response
|
||||
"""
|
||||
# Generate the prompt and log it
|
||||
prompt_result = dialectic_prompt(
|
||||
query,
|
||||
working_representation,
|
||||
recent_conversation_history,
|
||||
additional_context,
|
||||
peer_name,
|
||||
target_name,
|
||||
# Generate the prompt
|
||||
prompt_content = dialectic_prompt(
|
||||
query=query,
|
||||
working_representation=working_representation,
|
||||
recent_conversation_history=recent_conversation_history,
|
||||
additional_context=additional_context,
|
||||
peer_name=peer_name,
|
||||
target_name=target_name,
|
||||
)
|
||||
|
||||
# Pretty print the prompt content
|
||||
if len(prompt_result) > 0:
|
||||
# Extract content from the first BaseMessageParam
|
||||
prompt_content = prompt_result[0].content
|
||||
else:
|
||||
prompt_content = str(prompt_result)
|
||||
|
||||
logger.debug("=== DIALECTIC PROMPT (STREAM) ===")
|
||||
logger.debug(prompt_content)
|
||||
logger.debug("=== END DIALECTIC PROMPT ===")
|
||||
|
||||
return prompt_result
|
||||
# Make streaming LLM call
|
||||
stream = await direct_llm_call(
|
||||
prompt=prompt_content,
|
||||
provider=settings.DIALECTIC.PROVIDER,
|
||||
model=settings.DIALECTIC.MODEL,
|
||||
max_tokens=settings.DIALECTIC.MAX_OUTPUT_TOKENS,
|
||||
thinking_budget_tokens=settings.DIALECTIC.THINKING_BUDGET_TOKENS
|
||||
if settings.DIALECTIC.PROVIDER == "anthropic"
|
||||
else None,
|
||||
stream=True,
|
||||
track_name="Dialectic Stream",
|
||||
)
|
||||
|
||||
async for chunk in stream:
|
||||
yield chunk
|
||||
|
||||
|
||||
async def chat(
|
||||
|
|
@ -155,7 +146,7 @@ async def chat(
|
|||
query: str,
|
||||
*,
|
||||
stream: bool = False,
|
||||
) -> Stream | str:
|
||||
) -> AsyncGenerator[str, None] | str:
|
||||
"""
|
||||
Chat with the Dialectic API that builds on-demand user representations.
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from langfuse.decorators import langfuse_context
|
|||
|
||||
from src.config import settings
|
||||
from src.models import Document
|
||||
from src.utils.clients import honcho_llm_call
|
||||
from src.utils.clients import create_retry_wrapper, direct_llm_call
|
||||
from src.utils.embedding_store import EmbeddingStore
|
||||
from src.utils.formatting import (
|
||||
format_premises_for_display,
|
||||
|
|
@ -228,13 +228,21 @@ def _format_observations(
|
|||
return "\n".join(parts).strip()
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.DIALECTIC.QUERY_GENERATION_PROVIDER,
|
||||
model=settings.DIALECTIC.QUERY_GENERATION_MODEL,
|
||||
response_model=SemanticQueries,
|
||||
enable_retry=True,
|
||||
retry_attempts=3,
|
||||
)
|
||||
async def generate_semantic_queries(query: str, target_peer_name: str):
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def generate_semantic_queries(
|
||||
query: str, target_peer_name: str
|
||||
) -> SemanticQueries:
|
||||
"""Generate semantic search queries for observation retrieval."""
|
||||
return query_generation_prompt(query, target_peer_name)
|
||||
prompt_content = query_generation_prompt(
|
||||
query=query, target_peer_name=target_peer_name
|
||||
)
|
||||
|
||||
response = await direct_llm_call(
|
||||
prompt=prompt_content,
|
||||
provider=settings.DIALECTIC.QUERY_GENERATION_PROVIDER,
|
||||
model=settings.DIALECTIC.QUERY_GENERATION_MODEL,
|
||||
response_model=SemanticQueries,
|
||||
track_name="Query Generation Call",
|
||||
)
|
||||
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -7,11 +7,9 @@ from fastapi import (
|
|||
Depends,
|
||||
Path,
|
||||
)
|
||||
from fastapi.exceptions import HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from fastapi_pagination import Page
|
||||
from fastapi_pagination.ext.sqlalchemy import apaginate
|
||||
from mirascope.llm import Stream
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src import crud, schemas
|
||||
|
|
@ -188,14 +186,14 @@ async def chat(
|
|||
query=options.query,
|
||||
stream=options.stream,
|
||||
)
|
||||
if isinstance(stream, Stream):
|
||||
async for chunk, _ in stream:
|
||||
yield chunk.content
|
||||
if isinstance(stream, AsyncGenerator):
|
||||
async for chunk in stream:
|
||||
yield chunk
|
||||
else:
|
||||
raise HTTPException(status_code=500, detail="Invalid stream type")
|
||||
yield str(stream)
|
||||
except Exception as e:
|
||||
logger.error(f"Error in stream: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e)) from e
|
||||
yield f"Error: {str(e)}"
|
||||
|
||||
return StreamingResponse(
|
||||
content=parse_stream(), media_type="text/event-stream", status_code=200
|
||||
|
|
|
|||
|
|
@ -1,21 +1,9 @@
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import (
|
||||
Any,
|
||||
Literal,
|
||||
ParamSpec,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
overload,
|
||||
runtime_checkable,
|
||||
)
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from anthropic import AsyncAnthropic
|
||||
from google import genai
|
||||
from groq import AsyncGroq
|
||||
from mirascope import llm
|
||||
from mirascope.core import ResponseModelConfigDict
|
||||
from mirascope.integrations.langfuse import with_langfuse
|
||||
from mirascope.llm import Stream
|
||||
from openai import AsyncOpenAI
|
||||
from pydantic import BaseModel
|
||||
from sentry_sdk.ai.monitoring import ai_track
|
||||
|
|
@ -61,297 +49,248 @@ for provider_name, provider_value in providers:
|
|||
if provider_value not in clients:
|
||||
raise ValueError(f"Missing client for {provider_name}: {provider_value}")
|
||||
|
||||
P = ParamSpec("P")
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
T_co = TypeVar("T_co", bound=BaseModel, covariant=True)
|
||||
F = TypeVar("F", bound=Callable[..., Any])
|
||||
|
||||
|
||||
# Define protocols for different return types
|
||||
@runtime_checkable
|
||||
class AsyncResponseModelCallable(Protocol[P, T_co]):
|
||||
async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> T_co: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SyncResponseModelCallable(Protocol[P, T_co]):
|
||||
def __call__(self, *args: P.args, **kwargs: P.kwargs) -> T_co: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AsyncStreamCallable(Protocol[P]):
|
||||
async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> Stream: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SyncStreamCallable(Protocol[P]):
|
||||
def __call__(self, *args: P.args, **kwargs: P.kwargs) -> Stream: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AsyncStringCallable(Protocol[P]):
|
||||
async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> str: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SyncStringCallable(Protocol[P]):
|
||||
def __call__(self, *args: P.args, **kwargs: P.kwargs) -> str: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AsyncCallResponseCallable(Protocol[P]):
|
||||
async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> llm.CallResponse: ...
|
||||
|
||||
|
||||
# Overload for stream=True with async function
|
||||
@overload
|
||||
def honcho_llm_call(
|
||||
*,
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: type[BaseModel] | None = None,
|
||||
async def direct_llm_call(
|
||||
prompt: str,
|
||||
provider: Providers,
|
||||
model: str,
|
||||
response_model: type[T] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: Literal[True],
|
||||
**extra_call_params: Any,
|
||||
) -> Callable[[Callable[P, Awaitable[Any]]], AsyncStreamCallable[P]]: ...
|
||||
|
||||
|
||||
# Overload for response_model with async function
|
||||
@overload
|
||||
def honcho_llm_call(
|
||||
*,
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: type[T],
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: Literal[False] = False,
|
||||
**extra_call_params: Any,
|
||||
) -> Callable[[Callable[P, Awaitable[Any]]], AsyncResponseModelCallable[P, T]]: ...
|
||||
|
||||
|
||||
# Overload for return_call_response=True with async function
|
||||
@overload
|
||||
def honcho_llm_call(
|
||||
*,
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: Literal[False] = False,
|
||||
return_call_response: Literal[True],
|
||||
**extra_call_params: Any,
|
||||
) -> Callable[[Callable[P, Awaitable[Any]]], AsyncCallResponseCallable[P]]: ...
|
||||
|
||||
|
||||
# Overload for no response_model with async function (string return)
|
||||
@overload
|
||||
def honcho_llm_call(
|
||||
*,
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: Literal[False] = False,
|
||||
return_call_response: Literal[False],
|
||||
**extra_call_params: Any,
|
||||
) -> Callable[[Callable[P, Awaitable[Any]]], AsyncStringCallable[P]]: ...
|
||||
|
||||
|
||||
# Generic overload for sync functions (fallback)
|
||||
@overload
|
||||
def honcho_llm_call(
|
||||
*,
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: type[BaseModel] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: bool = False,
|
||||
**extra_call_params: Any,
|
||||
) -> Callable[[Callable[P, Any]], Callable[P, Any]]: ...
|
||||
|
||||
|
||||
def honcho_llm_call(
|
||||
provider: Providers | None = None,
|
||||
model: str | None = None,
|
||||
track_name: str | None = None,
|
||||
response_model: type[BaseModel] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
enable_retry: bool = True,
|
||||
retry_attempts: int = 3,
|
||||
stream: bool = False,
|
||||
return_call_response: bool = False, # pyright: ignore
|
||||
**extra_call_params: Any,
|
||||
) -> Any:
|
||||
# track_name: str | None = None,
|
||||
) -> T | str | AsyncGenerator[str, None]:
|
||||
"""
|
||||
Consolidated decorator for LLM calls that handles provider-specific configurations.
|
||||
|
||||
This decorator automatically:
|
||||
- Handles both sync and async functions seamlessly
|
||||
- Applies retry logic with exponential backoff
|
||||
- Adds AI tracking for Sentry
|
||||
- Integrates with Langfuse for observability
|
||||
- Builds provider-specific call parameters
|
||||
- Handles client selection from the global clients dict
|
||||
Direct LLM call using native client libraries.
|
||||
|
||||
Args:
|
||||
provider: The LLM provider to use (e.g., "anthropic", "google", "openai")
|
||||
model: The model to use
|
||||
track_name: Name for AI tracking (e.g., "Critical Analysis Call")
|
||||
response_model: Optional Pydantic model for structured responses
|
||||
json_mode: Whether to enable JSON mode (for providers that support it)
|
||||
max_tokens: Maximum tokens for the response
|
||||
prompt: The prompt text
|
||||
provider: LLM provider to use
|
||||
model: Model name
|
||||
response_model: Pydantic model for structured responses
|
||||
json_mode: Enable JSON mode
|
||||
max_tokens: Maximum tokens for response
|
||||
thinking_budget_tokens: Budget for thinking tokens (Anthropic only)
|
||||
enable_retry: Whether to enable retry logic (default: True)
|
||||
retry_attempts: Number of retry attempts (default: 3)
|
||||
stream: Whether to enable streaming responses (default: False)
|
||||
_return_call_response: Whether to return the full CallResponse object (default: False)
|
||||
**extra_call_params: Additional provider-specific parameters
|
||||
stream: Enable streaming
|
||||
track_name: Name for AI tracking
|
||||
|
||||
Returns:
|
||||
A decorator that returns:
|
||||
- For async functions: Callable[P, Awaitable[T]] where T is Stream, response_model, CallResponse, or str
|
||||
- For sync functions: Callable[P, T] where T is Stream, response_model, CallResponse, or str
|
||||
|
||||
Note: Type annotations may be needed at the call site for proper type checking.
|
||||
|
||||
Example (async function):
|
||||
@honcho_llm_call(
|
||||
provider=settings.DERIVER.PROVIDER,
|
||||
model=settings.DERIVER.MODEL,
|
||||
track_name="Critical Analysis Call",
|
||||
response_model=ReasoningResponse,
|
||||
json_mode=True,
|
||||
max_tokens=settings.DERIVER.MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
async def analyze(context: str, query: str):
|
||||
return prompt_template(context, query)
|
||||
|
||||
Example (sync function):
|
||||
@honcho_llm_call(
|
||||
provider="openai",
|
||||
model="gpt-4",
|
||||
max_tokens=1000,
|
||||
)
|
||||
def generate_summary(text: str) -> str:
|
||||
return f"Summarize: {text}"
|
||||
|
||||
# Call synchronously
|
||||
result = generate_summary("Long text here...")
|
||||
Response model instance, string, or streaming generator
|
||||
"""
|
||||
client = clients[provider]
|
||||
|
||||
def decorator(func: Callable[..., Any]) -> Callable[..., Any]:
|
||||
# Handle special case for custom provider
|
||||
# Custom providers use OpenAI-compatible endpoints, so we resolve to "openai" for the provider name
|
||||
# but keep the original "custom" for client lookup
|
||||
resolved_provider = "openai" if provider == "custom" else provider
|
||||
# if track_name:
|
||||
# Wrap with AI tracking
|
||||
# from functools import wraps
|
||||
|
||||
# Build provider-specific call params
|
||||
call_params: dict[str, Any] = {}
|
||||
# def ai_track_decorator(func):
|
||||
# @wraps(func)
|
||||
# async def wrapper(*args, **kwargs):
|
||||
# return await ai_track(track_name)(func)(*args, **kwargs)
|
||||
#
|
||||
# return wrapper
|
||||
|
||||
if resolved_provider == "google":
|
||||
# Google uses 'config' parameter
|
||||
config: dict[str, Any] = {}
|
||||
if max_tokens:
|
||||
config["max_output_tokens"] = max_tokens
|
||||
# Handle custom provider (OpenAI-compatible)
|
||||
resolved_provider = "openai" if provider == "custom" else provider
|
||||
|
||||
if response_model:
|
||||
config["response_schema"] = response_model
|
||||
if resolved_provider == "google":
|
||||
return await _call_google(
|
||||
client=client, # pyright: ignore
|
||||
prompt=prompt,
|
||||
model=model,
|
||||
response_model=response_model,
|
||||
json_mode=json_mode,
|
||||
max_tokens=max_tokens,
|
||||
stream=stream,
|
||||
)
|
||||
elif resolved_provider == "anthropic":
|
||||
return await _call_anthropic(
|
||||
client=client, # pyright: ignore
|
||||
prompt=prompt,
|
||||
model=model,
|
||||
response_model=response_model,
|
||||
json_mode=json_mode,
|
||||
max_tokens=max_tokens,
|
||||
thinking_budget_tokens=thinking_budget_tokens,
|
||||
stream=stream,
|
||||
)
|
||||
else: # openai, groq
|
||||
return await _call_openai_compatible(
|
||||
client=client, # pyright: ignore
|
||||
prompt=prompt,
|
||||
model=model,
|
||||
response_model=response_model,
|
||||
json_mode=json_mode,
|
||||
max_tokens=max_tokens,
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
if json_mode:
|
||||
config["response_mime_type"] = "application/json"
|
||||
|
||||
if config:
|
||||
call_params["config"] = config
|
||||
elif resolved_provider == "anthropic":
|
||||
# Anthropic uses thinking params and max_tokens
|
||||
if thinking_budget_tokens:
|
||||
call_params["thinking"] = {
|
||||
"type": "enabled",
|
||||
"budget_tokens": thinking_budget_tokens,
|
||||
}
|
||||
if max_tokens:
|
||||
call_params["max_tokens"] = max_tokens
|
||||
else:
|
||||
# Other providers just use max_tokens
|
||||
if max_tokens:
|
||||
call_params["max_tokens"] = max_tokens
|
||||
async def _call_google(
|
||||
client: genai.Client,
|
||||
prompt: str,
|
||||
model: str,
|
||||
response_model: type[T] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
stream: bool = False,
|
||||
) -> T | str | AsyncGenerator[str, None]:
|
||||
"""Google Gemini API call."""
|
||||
config: dict[str, Any] = {}
|
||||
if max_tokens:
|
||||
config["max_output_tokens"] = max_tokens
|
||||
|
||||
# Merge with any extra call params
|
||||
# Remove return_call_response from extra_call_params --
|
||||
# that one is just for our type system.
|
||||
extra_call_params.pop("return_call_response", None)
|
||||
call_params.update(extra_call_params)
|
||||
if response_model:
|
||||
config["response_schema"] = response_model
|
||||
|
||||
if json_mode:
|
||||
config["response_mime_type"] = "application/json"
|
||||
|
||||
if stream:
|
||||
response = client.models.generate_content_stream(
|
||||
model=model,
|
||||
contents=prompt,
|
||||
config=config,
|
||||
)
|
||||
|
||||
async def stream_generator() -> AsyncGenerator[str, None]:
|
||||
async for chunk in response:
|
||||
if chunk.text:
|
||||
yield chunk.text
|
||||
|
||||
return stream_generator()
|
||||
else:
|
||||
response = client.models.generate_content(
|
||||
model=model,
|
||||
contents=prompt,
|
||||
config=config,
|
||||
)
|
||||
|
||||
# Build kwargs for llm.call
|
||||
llm_kwargs: dict[str, Any] = {}
|
||||
if resolved_provider and provider:
|
||||
llm_kwargs["provider"] = resolved_provider
|
||||
llm_kwargs["client"] = clients[
|
||||
provider
|
||||
] # Use original provider for client lookup
|
||||
if model:
|
||||
llm_kwargs["model"] = model
|
||||
if response_model:
|
||||
# https://mirascope.com/docs/mirascope/learn/provider-specific/openai#response-models
|
||||
if resolved_provider == "openai":
|
||||
response_model.model_config = ResponseModelConfigDict(strict=True)
|
||||
return response_model.model_validate_json(response.text)
|
||||
elif response.text is None:
|
||||
return ""
|
||||
else:
|
||||
return response.text
|
||||
|
||||
llm_kwargs["response_model"] = response_model
|
||||
if json_mode:
|
||||
llm_kwargs["json_mode"] = json_mode
|
||||
if stream:
|
||||
llm_kwargs["stream"] = stream
|
||||
if call_params:
|
||||
llm_kwargs["call_params"] = call_params
|
||||
|
||||
# Apply decorators in order
|
||||
decorated: Any = func
|
||||
async def _call_anthropic(
|
||||
client: AsyncAnthropic,
|
||||
prompt: str,
|
||||
model: str,
|
||||
response_model: type[T] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
thinking_budget_tokens: int | None = None,
|
||||
stream: bool = False,
|
||||
) -> T | str | AsyncGenerator[str, None]:
|
||||
"""Anthropic Claude API call."""
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
|
||||
# Apply llm.call
|
||||
decorated = llm.call(**llm_kwargs)(decorated) # pyright: ignore
|
||||
call_params = {}
|
||||
if max_tokens:
|
||||
call_params["max_tokens"] = max_tokens
|
||||
|
||||
# Apply langfuse if enabled
|
||||
if settings.LANGFUSE_PUBLIC_KEY:
|
||||
decorated = with_langfuse()(decorated) # pyright: ignore
|
||||
if thinking_budget_tokens:
|
||||
call_params["thinking"] = {
|
||||
"type": "enabled",
|
||||
"budget_tokens": thinking_budget_tokens,
|
||||
}
|
||||
|
||||
# Apply AI tracking if name provided
|
||||
if track_name:
|
||||
decorated = ai_track(track_name)(decorated)
|
||||
if response_model or json_mode:
|
||||
call_params["response_format"] = {"type": "json_object"}
|
||||
|
||||
# Apply retry logic if enabled
|
||||
if enable_retry:
|
||||
decorated = retry( # pyright: ignore
|
||||
stop=stop_after_attempt(retry_attempts),
|
||||
wait=wait_exponential(multiplier=1, min=4, max=10),
|
||||
)(decorated) # pyright: ignore
|
||||
if stream:
|
||||
response = await client.messages.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
stream=True,
|
||||
**call_params,
|
||||
)
|
||||
|
||||
return decorated # pyright: ignore
|
||||
async def stream_generator():
|
||||
async for chunk in response:
|
||||
if chunk.type == "content_block_delta" and chunk.delta.text:
|
||||
yield chunk.delta.text
|
||||
|
||||
return decorator
|
||||
return stream_generator()
|
||||
else:
|
||||
response = await client.messages.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
**call_params,
|
||||
)
|
||||
|
||||
content = response.content[0].text
|
||||
|
||||
if response_model:
|
||||
return response_model.model_validate_json(content)
|
||||
else:
|
||||
return content
|
||||
|
||||
|
||||
async def _call_openai_compatible(
|
||||
client: AsyncOpenAI,
|
||||
prompt: str,
|
||||
model: str,
|
||||
response_model: type[T] | None = None,
|
||||
json_mode: bool = False,
|
||||
max_tokens: int | None = None,
|
||||
stream: bool = False,
|
||||
) -> T | str | AsyncGenerator[str, None]:
|
||||
"""OpenAI-compatible API call (OpenAI, Groq, custom providers)."""
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
|
||||
call_params: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
}
|
||||
|
||||
if max_tokens:
|
||||
call_params["max_tokens"] = max_tokens
|
||||
|
||||
if response_model:
|
||||
call_params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": response_model.__name__,
|
||||
"schema": response_model.model_json_schema(),
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
elif json_mode:
|
||||
call_params["response_format"] = {"type": "json_object"}
|
||||
|
||||
if stream:
|
||||
call_params["stream"] = True
|
||||
response = await client.chat.completions.create(**call_params)
|
||||
|
||||
async def stream_generator():
|
||||
async for chunk in response:
|
||||
if chunk.choices and chunk.choices[0].delta.content:
|
||||
yield chunk.choices[0].delta.content
|
||||
|
||||
return stream_generator()
|
||||
else:
|
||||
response = await client.chat.completions.create(**call_params)
|
||||
content = response.choices[0].message.content
|
||||
|
||||
if response_model:
|
||||
return response_model.model_validate_json(content or "")
|
||||
else:
|
||||
return content or ""
|
||||
|
||||
|
||||
def create_retry_wrapper(max_attempts: int = 3):
|
||||
"""Create retry decorator with exponential backoff."""
|
||||
return retry(
|
||||
stop=stop_after_attempt(max_attempts),
|
||||
wait=wait_exponential(multiplier=1, min=4, max=10),
|
||||
)
|
||||
|
||||
|
||||
# Keep the old honcho_llm_call for now, but mark as deprecated
|
||||
# We'll remove it after all usages are migrated
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import logging
|
|||
import time
|
||||
from enum import Enum
|
||||
|
||||
from mirascope import llm
|
||||
from sqlalchemy import update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -12,7 +11,7 @@ from typing_extensions import TypedDict
|
|||
from src.config import settings
|
||||
from src.dependencies import tracked_db
|
||||
from src.exceptions import ResourceNotFoundException
|
||||
from src.utils.clients import honcho_llm_call
|
||||
from src.utils.clients import create_retry_wrapper, direct_llm_call
|
||||
from src.utils.logging import accumulate_metric
|
||||
|
||||
from .. import crud, models
|
||||
|
|
@ -63,17 +62,12 @@ class SummaryType(Enum):
|
|||
LONG = "honcho_chat_summary_long"
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.SUMMARY.PROVIDER,
|
||||
model=settings.SUMMARY.MODEL,
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_SHORT,
|
||||
return_call_response=True,
|
||||
)
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def create_short_summary(
|
||||
messages: list[models.Message],
|
||||
input_tokens: int,
|
||||
previous_summary: str | None = None,
|
||||
):
|
||||
) -> str:
|
||||
# input_tokens indicates how many tokens the message list + previous summary take up
|
||||
# we want to optimize short summaries to be smaller than the actual content being summarized
|
||||
# so we ask the agent to produce a word count roughly equal to either the input, or the max
|
||||
|
|
@ -112,16 +106,116 @@ Produce as thorough a summary as possible in {output_words} words or less.
|
|||
"""
|
||||
|
||||
|
||||
@honcho_llm_call(
|
||||
provider=settings.SUMMARY.PROVIDER,
|
||||
model=settings.SUMMARY.MODEL,
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_LONG,
|
||||
return_call_response=True,
|
||||
)
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def create_short_summary(
|
||||
messages: list[models.Message],
|
||||
input_tokens: int,
|
||||
previous_summary: str | None = None,
|
||||
) -> str:
|
||||
# input_tokens indicates how many tokens the message list + previous summary take up
|
||||
# we want to optimize short summaries to be smaller than the actual content being summarized
|
||||
# so we ask the agent to produce a word count roughly equal to either the input, or the max
|
||||
# size if the input is larger. the word/token ratio is roughly 4:3 so we multiply by 0.75.
|
||||
# LLMs *seem* to respond better to getting asked for a word count but should workshop this.
|
||||
output_words = int(min(input_tokens, settings.SUMMARY.MAX_TOKENS_SHORT) * 0.75)
|
||||
|
||||
if previous_summary:
|
||||
previous_summary_text = previous_summary
|
||||
else:
|
||||
previous_summary_text = "There is no previous summary -- the messages are the beginning of the conversation."
|
||||
|
||||
prompt = f"""
|
||||
You are a system that summarizes parts of a conversation to create a concise and accurate summary. Focus on capturing:
|
||||
|
||||
1. Key facts and information shared (**Capture as many explicit facts as possible**)
|
||||
2. User preferences, opinions, and questions
|
||||
3. Important context and requests
|
||||
4. Core topics discussed
|
||||
|
||||
If there is a previous summary, ALWAYS make your new summary inclusive of both it and the new messages, therefore capturing the ENTIRE conversation. Prioritize key facts across the entire conversation.
|
||||
|
||||
Provide a concise, factual summary that captures the essence of the conversation. Your summary should be detailed enough to serve as context for future messages, but brief enough to be helpful. Prefer a thorough chronological narrative over a list of bullet points.
|
||||
|
||||
Return only the summary without any explanation or meta-commentary.
|
||||
|
||||
<previous_summary>
|
||||
{previous_summary_text}
|
||||
</previous_summary>
|
||||
|
||||
<conversation>
|
||||
{_format_messages(messages)}
|
||||
</conversation>
|
||||
|
||||
Produce as thorough a summary as possible in {output_words} words or less.
|
||||
"""
|
||||
|
||||
response = await direct_llm_call(
|
||||
prompt=prompt,
|
||||
provider=settings.SUMMARY.PROVIDER,
|
||||
model=settings.SUMMARY.MODEL,
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_SHORT,
|
||||
track_name="Short Summary Call",
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@create_retry_wrapper(max_attempts=3)
|
||||
async def create_long_summary(
|
||||
messages: list[models.Message],
|
||||
previous_summary: str | None = None,
|
||||
):
|
||||
) -> str:
|
||||
# the word/token ratio is roughly 4:3 so we multiply by 0.75.
|
||||
# LLMs *seem* to respond better to getting asked for a word count but should workshop this.
|
||||
output_words = int(settings.SUMMARY.MAX_TOKENS_LONG * 0.75)
|
||||
|
||||
if previous_summary:
|
||||
previous_summary_text = previous_summary
|
||||
else:
|
||||
previous_summary_text = "There is no previous summary -- the messages are the beginning of the conversation."
|
||||
|
||||
prompt = f"""
|
||||
You are a system that creates thorough, comprehensive summaries of conversations. Focus on capturing:
|
||||
|
||||
1. Key facts and information shared (**Capture as many explicit facts as possible**)
|
||||
2. User preferences, opinions, and questions
|
||||
3. Important context and requests
|
||||
4. Core topics discussed in detail
|
||||
5. User's apparent emotional state and personality traits
|
||||
6. Important themes and patterns across the conversation
|
||||
|
||||
If there is a previous summary, ALWAYS make your new summary inclusive of both it and the new messages, therefore capturing the ENTIRE conversation. Prioritize key facts across the entire conversation.
|
||||
|
||||
Provide a comprehensive, detailed summary that thoroughly captures the entire conversation. Aim for completeness over brevity, but keep it focused on actionable insights and facts.
|
||||
|
||||
Return only the summary without any explanation or meta-commentary.
|
||||
|
||||
<previous_summary>
|
||||
{previous_summary_text}
|
||||
</previous_summary>
|
||||
|
||||
<conversation>
|
||||
{_format_messages(messages)}
|
||||
</conversation>
|
||||
|
||||
Create a comprehensive summary in approximately {output_words} words.
|
||||
"""
|
||||
|
||||
response = await direct_llm_call(
|
||||
prompt=prompt,
|
||||
provider=settings.SUMMARY.PROVIDER,
|
||||
model=settings.SUMMARY.MODEL,
|
||||
max_tokens=settings.SUMMARY.MAX_TOKENS_LONG,
|
||||
track_name="Long Summary Call",
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
async def create_long_summary(
|
||||
messages: list[models.Message],
|
||||
previous_summary: str | None = None,
|
||||
) -> str:
|
||||
# the word/token ratio is roughly 4:3 so we multiply by 0.75.
|
||||
# LLMs *seem* to respond better to getting asked for a word count but should workshop this.
|
||||
output_words = int(settings.SUMMARY.MAX_TOKENS_LONG * 0.75)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,3 @@
|
|||
from typing import Literal
|
||||
|
||||
from mirascope import Provider
|
||||
|
||||
Providers = Provider | Literal["custom"]
|
||||
Providers = Literal["anthropic", "openai", "google", "groq", "custom"]
|
||||
|
|
|
|||
|
|
@ -520,12 +520,13 @@ def test_get_peers_with_complex_filter(
|
|||
"filter": {
|
||||
"AND": [
|
||||
{"metadata": {"type": "test"}},
|
||||
{"metadata": {"index": {"gte": 1}}},
|
||||
{"metadata": {"index": {"gte": 2}}},
|
||||
]
|
||||
}
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert len(response.json()["items"]) == 3
|
||||
data = response.json()
|
||||
assert "items" in data
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue