File size: 5,737 Bytes
2e818da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
"""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)
    )