Gilbra's picture
on enleve claude
9947278
Raw
History Blame Contribute Delete
11.4 kB
import os
import re
import json
import logging
from typing import Any, List, Dict, Optional
from openai import OpenAI, AzureOpenAI
# =====================================================
# LOGGING SETUP
# =====================================================
_LOG_LEVEL = os.getenv("LOG_LEVEL", "INFO").upper()
logging.basicConfig(
level=getattr(logging, _LOG_LEVEL, logging.INFO),
format="%(asctime)s | %(levelname)s | %(message)s"
)
logger = logging.getLogger("TLC_AGENT_UTILS")
def _log_info(message: str) -> None:
logger.info(message)
def _log_warning(message: str) -> None:
logger.warning(message)
def _log_error(message: str) -> None:
logger.error(message)
# =====================================================
# JSON HELPERS
# =====================================================
def clean_json_text(text: str) -> str:
"""
Retire les fences markdown et nettoie le texte JSON.
"""
if not text:
return ""
text = str(text).strip()
# retire fences markdown
text = re.sub(r"^```json\s*", "", text, flags=re.IGNORECASE)
text = re.sub(r"^```\s*", "", text)
text = re.sub(r"\s*```$", "", text)
return text.strip()
def safe_json_loads(text: Any):
"""
Parse JSON robuste pour les réponses LLM.
Retourne dict/list ou None.
"""
if text is None:
return None
if isinstance(text, (dict, list)):
return text
text = str(text).strip()
if not text:
return None
_log_info(f"safe_json_loads -> raw preview: {text[:300]}")
text = clean_json_text(text)
# 1) tableau JSON
array_match = re.search(r"(\[[\s\S]*\])", text)
if array_match:
candidate = array_match.group(1)
try:
parsed = json.loads(candidate)
_log_info("safe_json_loads -> parsed as array")
return parsed
except Exception as e:
_log_warning(f"safe_json_loads -> array parse failed: {e}")
# 2) objet JSON
object_match = re.search(r"(\{[\s\S]*\})", text)
if object_match:
candidate = object_match.group(1)
try:
parsed = json.loads(candidate)
_log_info("safe_json_loads -> parsed as object")
return parsed
except Exception as e:
_log_warning(f"safe_json_loads -> object parse failed: {e}")
# 3) parse direct
try:
parsed = json.loads(text)
_log_info("safe_json_loads -> parsed directly")
return parsed
except Exception as e:
_log_error(f"safe_json_loads -> final parse failed: {e}")
return None
# =====================================================
# LLM CALLS
# =====================================================
def call_deepseek(
prompt: str,
model: str = "deepseek-chat",
max_tokens: int = 2000
) -> str:
"""
DeepSeek via API compatible OpenAI.
"""
api_key = os.getenv("DEEPSEEK_API_KEY")
if not api_key:
_log_error("DEEPSEEK_API_KEY manquante")
return "ERREUR: DEEPSEEK_API_KEY manquante"
try:
_log_info(f"Calling DeepSeek model={model}, max_tokens={max_tokens}")
client = OpenAI(
api_key=api_key,
base_url="https://api.deepseek.com"
)
response = client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": prompt}],
max_tokens=max_tokens,
temperature=0.3
)
content = response.choices[0].message.content or ""
_log_info(f"DeepSeek response length: {len(content)}")
return content.strip()
except Exception as e:
_log_error(f"ERREUR DEEPSEEK: {e}")
return f"ERREUR DEEPSEEK: {e}"
def call_groq(
prompt: str,
model: str = "llama-3.3-70b-versatile",
max_tokens: int = 2000
) -> str:
"""
Groq via OpenAI-compatible endpoint.
Évite les soucis du SDK Groq et les conflits de version.
"""
api_key = os.getenv("GROQ_API_KEY")
if not api_key:
_log_error("GROQ_API_KEY manquante")
return "ERREUR: GROQ_API_KEY manquante"
try:
_log_info(f"Calling Groq model={model}, max_tokens={max_tokens}")
client = OpenAI(
api_key=api_key,
base_url="https://api.groq.com/openai/v1"
)
response = client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": prompt}],
max_tokens=max_tokens,
temperature=0.3
)
content = response.choices[0].message.content or ""
_log_info(f"Groq response length: {len(content)}")
return content.strip()
except Exception as e:
_log_error(f"ERREUR GROQ: {e}")
return f"ERREUR GROQ: {e}"
def call_azure_openai(
prompt: str,
max_tokens: int = 2000
) -> str:
"""
Azure OpenAI.
OPENAI_MODEL doit correspondre au nom du déploiement Azure.
"""
api_key = os.getenv("AZUREOPENAI_API_KEY")
endpoint = os.getenv("AZUREOPENAI_API_ENDPOINT")
api_version = os.getenv("AZUREOPENAI_API_VERSION")
deployment = os.getenv("OPENAI_MODEL", "gpt-4o-mini")
if not api_key or not endpoint or not api_version:
_log_error("Azure OpenAI credentials manquantes")
return "ERREUR: Azure OpenAI credentials manquantes"
try:
_log_info(
f"Calling Azure OpenAI deployment={deployment}, max_tokens={max_tokens}"
)
client = AzureOpenAI(
api_key=api_key,
api_version=api_version,
azure_endpoint=endpoint
)
response = client.chat.completions.create(
model=deployment,
messages=[{"role": "user", "content": prompt}],
max_tokens=max_tokens,
temperature=0.3
)
content = response.choices[0].message.content or ""
_log_info(f"Azure OpenAI response length: {len(content)}")
return content.strip()
except Exception as e:
_log_error(f"ERREUR AZURE OPENAI: {e}")
return f"ERREUR AZURE OPENAI: {e}"
def call_llm(
prompt: str,
provider: str = "azure",
max_tokens: int = 2000
) -> str:
"""
Routeur principal.
"""
provider = (provider or "").lower().strip()
_log_info(f"call_llm -> provider={provider}")
if provider == "azure":
return call_azure_openai(prompt, max_tokens=max_tokens)
if provider == "groq":
return call_groq(prompt, max_tokens=max_tokens)
if provider == "deepseek":
return call_deepseek(prompt, max_tokens=max_tokens)
_log_error(f"Provider inconnu: {provider}")
return "ERREUR: provider inconnu"
def smart_call(prompt: str, max_tokens: int = 2000) -> str:
"""
Fallback automatique.
Ordre conseillé: Azure -> Groq -> DeepSeek.
"""
providers = ["azure", "groq", "deepseek"]
for provider in providers:
_log_info(f"smart_call -> trying {provider}")
result = call_llm(prompt, provider=provider, max_tokens=max_tokens)
if isinstance(result, str) and not result.startswith("ERREUR"):
_log_info(f"smart_call -> success with {provider}")
return result
_log_error("smart_call -> all providers failed")
return "ERREUR: tous les providers ont échoué"
# =====================================================
# LATEX EXTRACTION
# =====================================================
def extract_latex_blocks(text: str) -> List[str]:
"""
Extraction rapide des blocs LaTeX:
- $$...$$
- \[...\]
- $...$
"""
if not text:
return []
pattern = r'\$\$([^\$]+)\$\$|\\\[(.*?)\\\]|\$([^\$]+)\$'
matches = re.findall(pattern, text, re.DOTALL)
equations = []
for m in matches:
eq = m[0] or m[1] or m[2]
if eq and eq.strip():
equations.append(eq.strip())
unique = []
seen = set()
for eq in equations:
if eq not in seen:
seen.add(eq)
unique.append(eq)
_log_info(f"extract_latex_blocks -> found {len(unique)} equations")
return unique
def _normalize_llm_equation_item(item: Any) -> Optional[Dict[str, str]]:
"""
Normalise un item issu d'une réponse LLM en {"latex": "..."}.
"""
if item is None:
return None
if isinstance(item, str):
s = item.strip()
if s:
return {"latex": s}
return None
if isinstance(item, dict):
latex = item.get("latex") or item.get("equation") or item.get("expr")
if latex and str(latex).strip():
return {"latex": str(latex).strip()}
return None
s = str(item).strip()
if s:
return {"latex": s}
return None
def extract_equations_with_llm(
text: str,
provider: str = "azure"
) -> List[Dict[str, str]]:
"""
Extrait des équations mathématiques via LLM.
Retourne une liste de dicts: [{"latex": "..."}]
"""
prompt = f"""
Tu es un assistant scientifique.
Extrait uniquement les équations mathématiques présentes dans ce texte.
Retourne STRICTEMENT un JSON valide sous cette forme:
[
{{
"latex": "E = mc^2"
}}
]
Contraintes:
- Aucun texte hors JSON
- Pas de markdown
- Pas d'explication
Texte:
{text}
"""
_log_info(f"extract_equations_with_llm -> sending to provider={provider}")
response = call_llm(
prompt,
provider=provider,
max_tokens=2000
)
data = safe_json_loads(response)
if data is None:
_log_warning("extract_equations_with_llm -> parsing failed")
return []
normalized: List[Dict[str, str]] = []
if isinstance(data, list):
for item in data:
norm = _normalize_llm_equation_item(item)
if norm:
normalized.append(norm)
else:
norm = _normalize_llm_equation_item(data)
if norm:
normalized.append(norm)
_log_info(
f"extract_equations_with_llm -> parsed {len(normalized)} equations"
)
return normalized
# =====================================================
# IR HELPERS
# =====================================================
def validate_ir(ir: Any) -> bool:
if not isinstance(ir, dict):
return False
if "nodes" not in ir or "edges" not in ir:
return False
return True
def normalize_ir(ir: Any) -> Dict[str, Any]:
if not isinstance(ir, dict):
_log_warning("normalize_ir -> invalid IR, building fallback shell")
return {
"name": "Invalid IR",
"strategy": "fallback",
"nodes": [],
"edges": []
}
ir.setdefault("name", "Unnamed IR")
ir.setdefault("strategy", "fallback")
ir.setdefault("nodes", [])
ir.setdefault("edges", [])
return ir
def build_fallback_ir(equations: List[Any]) -> Dict[str, Any]:
nodes = []
for i, eq in enumerate(equations):
latex = eq.get("latex", "") if isinstance(eq, dict) else str(eq)
nodes.append({
"id": f"eq_{i}",
"type": "equation",
"latex": latex
})
_log_info(f"build_fallback_ir -> built IR with {len(nodes)} nodes")
return {
"name": "Fallback Variant 1",
"strategy": "fallback",
"nodes": nodes,
"edges": []
}