study-buddy / app /services /cognee_bootstrap.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
5.74 kB
"""Cognee runtime configuration for the local Cerebras/Gemma setup."""
from __future__ import annotations
import json
import os
import importlib
from typing import Any, Callable
from app.agents.cerebras_client import CerebrasClient
COGNEE_LLM_MODEL = "openai/gemma-4-31b"
COGNEE_CEREBRAS_BASE_URL = "https://api.cerebras.ai/v1"
COGNEE_INSTRUCTOR_MODE = "json_schema_mode"
COGNEE_SKIP_CONNECTION_TEST = "true"
COGNEE_OPENAI_ADAPTER_MODULE = (
"cognee.infrastructure.llm.structured_output_framework.litellm_instructor.llm.openai.adapter"
)
def _cognee_llm_args() -> dict[str, Any]:
"""LiteLLM args shared by all Cognee structured-output calls."""
return {
"temperature": 0,
"top_p": 1,
"seed": 0,
}
def configure_cognee_llm(
cognee_config: Any,
*,
clear_llm_client_cache: Callable[[], None] | None = None,
) -> dict[str, Any]:
"""Force Cognee's LiteLLM/instructor path onto Cerebras strict JSON mode.
StudyBuddy's own ``CerebrasClient`` already uses provider-native
``response_format={"type": "json_schema", ... "strict": true}``. Cognee
reaches Cerebras through LiteLLM + instructor instead, so we configure both
the Cognee config object and the environment it reads from before any Cognee
client can be cached.
"""
cerebras_key = os.environ.get("CEREBRAS_API_KEY", "")
llm_args = _cognee_llm_args()
config = {
"structured_output_framework": "instructor",
"llm_provider": "openai",
"llm_model": COGNEE_LLM_MODEL,
"llm_endpoint": COGNEE_CEREBRAS_BASE_URL,
"llm_api_key": cerebras_key,
"llm_instructor_mode": COGNEE_INSTRUCTOR_MODE,
"llm_temperature": 0.0,
"llm_max_completion_tokens": 16384,
"llm_args": llm_args,
}
os.environ["OPENAI_API_KEY"] = cerebras_key
os.environ["OPENAI_API_BASE"] = COGNEE_CEREBRAS_BASE_URL
os.environ["LLM_API_KEY"] = cerebras_key
os.environ["LLM_API_BASE"] = COGNEE_CEREBRAS_BASE_URL
os.environ["LLM_ENDPOINT"] = COGNEE_CEREBRAS_BASE_URL
os.environ["LLM_MODEL"] = COGNEE_LLM_MODEL
os.environ["LLM_PROVIDER"] = "openai"
os.environ["LLM_INSTRUCTOR_MODE"] = COGNEE_INSTRUCTOR_MODE
os.environ["LLM_TEMPERATURE"] = "0"
os.environ["LLM_ARGS"] = json.dumps(llm_args)
os.environ["COGNEE_SKIP_CONNECTION_TEST"] = COGNEE_SKIP_CONNECTION_TEST
cognee_config.set_llm_config(config)
patch_cognee_cerebras_structured_output()
if clear_llm_client_cache is not None:
clear_llm_client_cache()
return config
def patch_cognee_cerebras_structured_output() -> bool:
"""Route Cognee+Cerebras structured outputs through native strict JSON schema.
Cognee's default OpenAI adapter uses LiteLLM + instructor. With Cerebras/Gemma
this can produce fenced or schema-shaped JSON that fails Cognee's Pydantic
models during memify/cognify. ResearchMate's own Cerebras client already uses
provider-native strict JSON schema; this patch makes Cognee use the same path
for BaseModel structured outputs while preserving the original adapter for
plain string calls and non-Cerebras endpoints.
"""
module = importlib.import_module(COGNEE_OPENAI_ADAPTER_MODULE)
adapter_cls = getattr(module, "OpenAIAdapter")
if getattr(adapter_cls, "_researchmate_cerebras_patch", False):
return True
original = getattr(adapter_cls, "acreate_structured_output", None)
async def _native_cerebras_structured_output(
self: Any,
text_input: str,
system_prompt: str,
response_model: type[Any],
**kwargs: Any,
) -> Any:
if response_model is str:
if hasattr(self, "acreate_str_output"):
return await self.acreate_str_output(text_input, system_prompt, **kwargs)
if original is not None:
return await original(self, text_input, system_prompt, response_model, **kwargs)
model = str(getattr(self, "model", "") or COGNEE_LLM_MODEL)
endpoint = str(getattr(self, "endpoint", "") or "")
is_cerebras = COGNEE_CEREBRAS_BASE_URL in endpoint or "gemma" in model.lower()
if not is_cerebras and original is not None:
return await original(self, text_input, system_prompt, response_model, **kwargs)
merged_kwargs = {**(getattr(self, "llm_args", {}) or {}), **kwargs}
max_completion_tokens = getattr(self, "max_completion_tokens", None)
if max_completion_tokens and "max_completion_tokens" not in merged_kwargs:
merged_kwargs["max_completion_tokens"] = max_completion_tokens
native_model = model.removeprefix("openai/")
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": text_input},
]
api_key = getattr(self, "api_key", None) or None
try:
client = CerebrasClient(api_key=api_key)
except TypeError:
client = CerebrasClient()
return await _run_sync_structured_complete(
client=client,
messages=messages,
response_model=response_model,
model=native_model,
kwargs=merged_kwargs,
)
adapter_cls.acreate_structured_output = _native_cerebras_structured_output
adapter_cls._researchmate_cerebras_patch = True
return True
async def _run_sync_structured_complete(
*,
client: CerebrasClient,
messages: list[dict[str, Any]],
response_model: type[Any],
model: str,
kwargs: dict[str, Any],
) -> Any:
import asyncio
return await asyncio.to_thread(
lambda: client.structured_complete(messages, response_model, model=model, **kwargs)
)