"""
K2 Think Engine - Modèle d'IA Principal Unique
Hackathon: K2 Think API est le seul moteur d'IA du projet
"""
from typing import List, Dict, Any, Optional
from datetime import datetime
import json
import os
import re
import traceback
import ast
from app.models.schemas import (
ScientificDocument, AnalysisResult, AnalysisRequest, AuditLog,
ComparativeAnalysis, ExperimentalProtocol, ResearchGap, CounterHypothesis,
ExperimentalStep, ExperimentalVariable, DocumentType
)
from app.reasoning.k2_client import K2ThinkClient
from app.core.settings import settings
from app.core.logging import logger
from langchain_openai import ChatOpenAI
from langchain.schema import HumanMessage
class K2ThinkEngine:
"""
Moteur K2 Think - IA Principal Unique
Toutes les analyses passent par l'API K2 Think exclusivement
"""
def __init__(self):
if not settings.K2_THINK_API_KEY:
raise ValueError("K2_THINK_API_KEY must be configured in .env")
self.k2_client = K2ThinkClient(
api_key=settings.K2_THINK_API_KEY,
api_url=settings.K2_THINK_API_URL
)
logger.info(f"K2 Think Engine initialized - UNIQUE AI MODEL for this hackathon")
logger.info(f" API URL: {settings.K2_THINK_API_URL}")
self.audit_logs: List[AuditLog] = []
self.reasoning_trace: List[Dict[str, Any]] = []
from app.services.memory_service import MemoryService
self.memory_service = MemoryService()
async def process_analysis_request(
self,
request: AnalysisRequest
) -> AnalysisResult:
"""
Processus complet d'analyse utilisant LangChain + K2 Think API
"""
request_id = f"analysis_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
self.reasoning_trace = []
self.audit_logs = []
try:
logger.info(f"K2 Think Analysis (LangChain Orchestrated) Start: {request_id}")
# 1. Préparation du contexte documentaire
context_parts = []
for doc in request.documents:
snippet = doc.content[:6000] # Limite pour économiser les tokens
first_author = doc.authors[0].split()[-1] if doc.authors else "Unknown"
year = "n.d."
citation_key = f"({first_author}, {year})"
context_parts.append(f"--- DOCUMENT: {doc.title} | KEY: {citation_key} ---\n{snippet}")
context = "\n\n".join(context_parts)
# 2. Prompt (Optimisé pour K2-Think-v2)
instruction_prompt = f"""You are a senior scientific investigator. Analyze the provided research documents and produce a detailed comparative analysis.
[DOCUMENTS]
{context}
[YOUR TASK]
1. Synthesize findings across all documents.
2. Identify divergences, contradictions, AND AT LEAST 2 RESEARCH GAPS OR OPPORTUNITIES.
3. Propose a new experimental protocol.
[FORMATTING RULES]
- Output ONLY a valid JSON object.
- NO preamble, NO explanations before or after JSON.
- KEEP REASONING CONCISE: Focus on direct analysis to stay within processing time limits.
- NO single quotes in the JSON keys or values.
- Use valid citations e.g. (Author, Year).
[JSON SCHEMA]
{{
"comparative_analysis": {{
"document_ids": {json.dumps([doc.id for doc in request.documents])},
"divergences": [
{{ "variable": "name", "finding_a": "...", "finding_b": "...", "impact": "..." }}
],
"contradictions": [
{{ "topic": "name", "conflict": "...", "resolution_path": "..." }}
],
"common_findings": ["Finding 1", "Finding 2"],
"confidence_score": 0.9
}},
"research_gaps": [
{{ "description": "Gap description", "importance_score": 0.8, "related_variables": [] }}
],
"counter_hypotheses": [
{{ "hypothesis": "...", "rationale": "...", "potential_bias": "...", "validation_experiment": "...", "confidence_against": 0.5 }}
],
"proposed_protocol": {{
"title": "...",
"objective": "...",
"hypothesis": "...",
"expected_outcomes": "...",
"variables": [
{{ "name": "...", "type": "independent", "measurement_method": "..." }}
],
"risk_assessment": {{ "overall_risk": "low" }},
"steps": [
{{ "description": "...", "duration_hours": 1, "materials": [], "critical_parameters": [] }}
]
}},
"strategic_recommendations": [],
"reasoning_summary": "Extensive 200+ word technical summary of findings",
"reasoning_trace": "Internal logic summary",
"confidence_overall": 0.95
}}
[FINAL INSTRUCTION]
YOU ARE STRICTLY FORBIDDEN FROM EXPLAINING YOUR REASONING.
YOU MUST RETURN ONLY THE RAW JSON OBJECT.
START YOUR RESPONSE DIRECTLY WITH {{ AND END WITH }}.
DO NOT USE TAGS. DO NOT CONVERSE.
"""
# 3. Appel au modèle (avec Retry Adaptatif en cas de Timeout)
chat_config = {
"model": "MBZUAI-IFM/K2-Think-v2",
"openai_api_key": settings.K2_THINK_API_KEY,
"openai_api_base": settings.K2_THINK_API_URL,
"max_tokens": 6000,
"timeout": 95,
"max_retries": 0
}
logger.info("Sending request to K2 Think...")
raw_content = ""
try:
chat = ChatOpenAI(**chat_config)
response = await chat.ainvoke([HumanMessage(content=instruction_prompt)])
raw_content = response.content
except Exception as e:
is_timeout = (
"524" in str(e)
or "timeout" in str(e).lower()
or type(e).__name__ in ["APITimeoutError", "Timeout", "ReadTimeout", "TimeoutError"]
)
if is_timeout:
logger.warning("K2 API Timeout detected. Retrying with very reduced context...")
# Maintain max_tokens to prevent JSON truncation, just rely on reduced context
chat_config["max_tokens"] = 5000
chat_config["timeout"] = 118
# Reduce the context even more aggressively to speed up generation
emergency_context = "\n\n".join([f"--- DOC: {d.title} ---\n{d.content[:1500]}" for d in request.documents])
emergency_prompt = instruction_prompt.replace(context, emergency_context)
chat = ChatOpenAI(**chat_config)
response = await chat.ainvoke([HumanMessage(content=emergency_prompt)])
raw_content = response.content
else:
raise e
logger.info(f"Raw K2 response length: {len(raw_content)}")
# 4. Extraction du JSON (Méthode robuste)
clean_json = ""
k2_analysis = None
# Nettoyage des balises de pensée
processed_content = raw_content
if "" in processed_content:
processed_content = processed_content.split("")[-1].strip()
elif "" in processed_content:
processed_content = re.sub(r'.*', '', processed_content, flags=re.DOTALL)
# Fonction utilitaire pour réparer le JSON
def repair_json(text):
repaired = text.strip()
# If it doesn't start with '{', find the first '{'
first_brace = repaired.find('{')
if first_brace != -1:
repaired = repaired[first_brace:]
if not repaired:
return repaired
# Fix unclosed quotes
in_string = False
escape = False
last_good_pos = len(repaired)
for i, char in enumerate(repaired):
if escape:
escape = False
continue
if char == '\\':
escape = True
continue
if char == '"':
in_string = not in_string
if not in_string:
last_good_pos = i + 1
if in_string:
repaired = repaired[:last_good_pos].strip()
# Remove trailing incomplete keys
# e.g., ',"key"' or ',"key":'
for suffix in ['', '}', ']']:
escaped_suffix = re.escape(suffix)
repaired = re.sub(r',\s*"[^"]*"\s*:\s*' + escaped_suffix + r'$', suffix, repaired)
repaired = re.sub(r',\s*"[^"]*"\s*' + escaped_suffix + r'$', suffix, repaired)
repaired = re.sub(r'\{\s*"[^"]*"\s*:\s*' + escaped_suffix + r'$', '{' + suffix, repaired)
repaired = re.sub(r'\{\s*"[^"]*"\s*' + escaped_suffix + r'$', '{' + suffix, repaired)
# Strip trailing punctuation/spaces
repaired = re.sub(r'[\s,:+]+$', '', repaired)
# Balance braces and brackets
open_braces = repaired.count('{')
close_braces = repaired.count('}')
open_brackets = repaired.count('[')
close_brackets = repaired.count(']')
while open_brackets > close_brackets:
repaired += ']'
close_brackets += 1
while open_braces > close_braces:
repaired += '}'
close_braces += 1
return repaired
candidates = []
# Stratégie 0 : Extraction par blocs markdown
for block in re.findall(r'```(?:json)?\s*(.*?)\s*```', processed_content, re.DOTALL | re.IGNORECASE):
b_start = block.find('{')
b_end = block.rfind('}')
if b_start != -1 and b_end != -1 and b_end > b_start:
candidates.append(block[b_start:b_end + 1])
else:
candidates.append(block)
# Stratégie 1 : Recherche par comptage de parenthèses (très robuste)
positions = [m.start() for m in re.finditer(r'\{', processed_content)]
for start_idx in positions:
brace_count = 0
in_string = False
escape = False
end_idx = -1
for i in range(start_idx, len(processed_content)):
char = processed_content[i]
if escape:
escape = False
continue
if char == '\\':
escape = True
continue
if char == '"':
in_string = not in_string
continue
if not in_string:
if char == '{':
brace_count += 1
elif char == '}':
brace_count -= 1
if brace_count == 0:
end_idx = i
break
if end_idx != -1:
candidates.append(processed_content[start_idx:end_idx + 1])
else:
candidates.append(processed_content[start_idx:]) # Truncated fallback
# Stratégie 2 : Blocs d'accolades globaux
start_idx = processed_content.find('{')
end_idx = processed_content.rfind('}')
if start_idx != -1 and end_idx != -1 and end_idx > start_idx:
candidates.append(processed_content[start_idx:end_idx + 1])
# Stratégie 3 : Recherche de "comparative_analysis"
first_idx = processed_content.find('"comparative_analysis"')
if first_idx != -1:
start_idx = processed_content.rfind('{', 0, first_idx)
if start_idx != -1:
candidates.append(processed_content[start_idx:])
# Prioritize candidates containing key schema terms
schema_keywords = ["comparative_analysis", "proposed_protocol", "reasoning_summary", "research_gaps"]
# Sort candidates by:
# 1. Matches at least one keyword (Boolean)
# 2. Length (longer is better for completeness)
def candidate_key(c):
has_keyword = any(kw in c for kw in schema_keywords)
return (1 if has_keyword else 0, len(c))
candidates.sort(key=candidate_key, reverse=True)
# 5. Parsing
import pathlib
debug_dir = pathlib.Path(__file__).parent.parent.parent
with open(debug_dir / "k2_debug_raw.txt", "w", encoding="utf-8") as f:
f.write(raw_content)
# Try parsing each candidate
for i, cand in enumerate(candidates):
if not cand.strip():
continue
with open(debug_dir / f"k2_debug_cand_{i}.txt", "w", encoding="utf-8") as f:
f.write(cand)
# 5.1 Direct parse
try:
k2_analysis = json.loads(cand)
clean_json = cand
break
except json.JSONDecodeError as e:
if "Extra data" in str(e) and hasattr(e, "pos"):
try:
k2_analysis = json.loads(cand[:e.pos].strip())
clean_json = cand[:e.pos].strip()
break
except Exception:
pass
# 5.2 Parse with repair
repaired_cand = repair_json(cand)
try:
k2_analysis = json.loads(repaired_cand)
clean_json = repaired_cand
break
except Exception:
# 5.3 Parse python literal dict fallback
try:
k2_analysis = ast.literal_eval(repaired_cand)
clean_json = repaired_cand
break
except Exception:
pass
if k2_analysis:
logger.info("JSON successfully extracted and parsed.")
# Recursive search helper to find nested objects/arrays
def find_key_recursive(data, target_key):
if isinstance(data, dict):
for k, v in data.items():
if k.lower() == target_key.lower():
return v
res = find_key_recursive(v, target_key)
if res is not None:
return res
elif isinstance(data, list):
for item in data:
res = find_key_recursive(item, target_key)
if res is not None:
return res
return None
# Unwrap if LLM wrapped everything in a single key like {"analysis": {...}}
if isinstance(k2_analysis, dict) and len(k2_analysis) == 1:
inner = list(k2_analysis.values())[0]
if isinstance(inner, dict):
k2_analysis = inner
# Unwrap if it's a list
if isinstance(k2_analysis, list) and len(k2_analysis) > 0 and isinstance(k2_analysis[0], dict):
k2_analysis = k2_analysis[0]
# Robust extraction: if critical keys are missing at root, search recursively
if isinstance(k2_analysis, dict):
if "comparative_analysis" not in k2_analysis:
nested_comp = find_key_recursive(k2_analysis, "comparative_analysis")
if nested_comp: k2_analysis["comparative_analysis"] = nested_comp
if "research_gaps" not in k2_analysis:
nested_gaps = find_key_recursive(k2_analysis, "research_gaps") or find_key_recursive(k2_analysis, "gaps")
if nested_gaps: k2_analysis["research_gaps"] = nested_gaps
if "counter_hypotheses" not in k2_analysis:
nested_hyp = find_key_recursive(k2_analysis, "counter_hypotheses")
if nested_hyp: k2_analysis["counter_hypotheses"] = nested_hyp
if "reasoning_summary" not in k2_analysis:
nested_summ = find_key_recursive(k2_analysis, "reasoning_summary") or find_key_recursive(k2_analysis, "summary")
if nested_summ: k2_analysis["reasoning_summary"] = nested_summ
# 6. Fallback en cas d'échec total de parsing
if not k2_analysis:
logger.error("All JSON parsing attempts failed. Creating technical fallback.")
self._log_reasoning("ERROR", "Parsing", f"Raw content snippet: {raw_content[:1000]}...")
# Salvage text to show to the user instead of a generic error
salvaged_text = processed_content.strip()
if not salvaged_text:
salvaged_text = raw_content.strip()
k2_analysis = {
"reasoning_summary": salvaged_text if len(salvaged_text) > 50 else f"Failed to extract structured data. Raw response: {raw_content[:1000]}",
"confidence_overall": 0.5,
"comparative_analysis": {
"document_ids": [doc.id for doc in request.documents],
"divergences": [],
"contradictions": [],
"common_findings": ["Partial analysis - structured data unavailable"],
"confidence_score": 0.5
},
"research_gaps": [],
"counter_hypotheses": [],
"proposed_protocol": {
"title": "Protocol Generation Failed",
"objective": "The AI generated text but could not format the experimental protocol correctly.",
"steps": []
},
"recommendations": ["Try reducing the number of documents", "The AI generated text but failed to format it as JSON"]
}
# 7. Conversion en objets schemas.py
comp_analysis = self._convert_k2_to_comparative_analysis(k2_analysis, request.documents)
hypotheses = self._convert_k2_to_counter_hypotheses(k2_analysis)
protocol = await self._convert_k2_to_protocol(k2_analysis)
# Ensure reasoning trace is populated for Audit Log
trace_val = k2_analysis.get("reasoning_trace")
if isinstance(trace_val, str):
self._log_reasoning("ANALYSIS", "K2 Synthesis", trace_val)
elif isinstance(trace_val, list):
for t in trace_val:
if isinstance(t, dict) and "reasoning" in t:
self.reasoning_trace.append(t)
else:
self._log_reasoning("ANALYSIS", "K2 Synthesis", str(t))
elif not self.reasoning_trace:
self._log_reasoning("ANALYSIS", "K2 Synthesis", "Analysis generated successfully.")
result = AnalysisResult(
request_id=request_id,
documents_analyzed=len(request.documents),
reasoning_summary=k2_analysis.get("reasoning_summary") or k2_analysis.get("summary") or k2_analysis.get("executive_summary") or "Analysis completed.",
comparative_analysis=comp_analysis,
research_gaps=comp_analysis.research_gaps,
counter_hypotheses=hypotheses,
proposed_protocol=protocol,
strategic_recommendations=k2_analysis.get("recommendations", k2_analysis.get("strategic_recommendations", [])),
reasoning_trace=self.reasoning_trace,
confidence_overall=k2_analysis.get("confidence_overall", 0.85)
)
# 8. Mémoire sémantique
if request.user_id:
try:
await self.memory_service.consolidate_analysis(
user_id=request.user_id,
project_id="auto_consolidation",
analysis_result=result
)
except Exception as mem_err:
logger.error(f"Memory consolidation failed: {mem_err}")
return result
except Exception as e:
logger.error(f"FATAL K2 Engine Error: {str(e)}")
traceback.print_exc()
raise e
def _convert_k2_to_comparative_analysis(
self,
k2_result: Dict[str, Any],
docs: List[ScientificDocument]
) -> ComparativeAnalysis:
raw_comp = k2_result.get("comparative_analysis", {})
if not isinstance(raw_comp, dict): raw_comp = {}
gaps = []
# Aggressive extraction for research_gaps
raw_gaps = (
k2_result.get("research_gaps") or
raw_comp.get("research_gaps") or
k2_result.get("gaps") or
raw_comp.get("gaps") or
k2_result.get("opportunities") or
[]
)
if isinstance(raw_gaps, dict):
raw_gaps = list(raw_gaps.values())[0] if raw_gaps else []
if not isinstance(raw_gaps, list):
raw_gaps = [raw_gaps]
for gap in raw_gaps:
if not gap: continue
if isinstance(gap, dict):
gaps.append(ResearchGap(
gap_description=gap.get("description", gap.get("gap_description", gap.get("name", "Research Gap Detected"))),
importance_score=float(gap.get("importance_score", gap.get("importance", 0.8))),
related_variables=gap.get("related_variables", gap.get("variables", [])),
suggested_investigation=gap.get("suggested_investigation", gap.get("investigation", "Investigation required")),
source_documents=[doc.id for doc in docs],
citations=gap.get("citations", [])
))
else:
gaps.append(ResearchGap(
gap_description=str(gap),
importance_score=0.8,
related_variables=[],
suggested_investigation="Investigation required",
source_documents=[doc.id for doc in docs]
))
return ComparativeAnalysis(
document_ids=[doc.id for doc in docs],
divergences=raw_comp.get("divergences", []),
contradictions=raw_comp.get("contradictions", []),
common_findings=raw_comp.get("common_findings", []),
research_gaps=gaps,
confidence_score=raw_comp.get("confidence_score", 0.8)
)
def _convert_k2_to_counter_hypotheses(
self,
k2_result: Dict[str, Any]
) -> List[CounterHypothesis]:
hypotheses = []
for h in k2_result.get("counter_hypotheses", []):
if isinstance(h, dict):
hypotheses.append(CounterHypothesis(
hypothesis=h.get("hypothesis", "Hypothesis"),
rationale=h.get("rationale", ""),
potential_bias=h.get("potential_bias", ""),
validation_experiment=h.get("validation_experiment", ""),
confidence_against=h.get("confidence_against", 0.5)
))
return hypotheses
async def _convert_k2_to_protocol(
self,
k2_result: Dict[str, Any]
) -> ExperimentalProtocol:
proto_data = k2_result.get("proposed_protocol", k2_result.get("protocol", {}))
if not isinstance(proto_data, dict): proto_data = {}
steps = []
for i, s in enumerate(proto_data.get("steps", []), 1):
if isinstance(s, dict):
steps.append(ExperimentalStep(
step_number=i,
description=s.get("description", f"Step {i}"),
duration_hours=float(s.get("duration_hours", 1)),
materials=s.get("materials", []),
critical_parameters=s.get("critical_parameters", [])
))
raw_vars = proto_data.get("variables", [])
valid_vars = []
for v in raw_vars:
if isinstance(v, dict):
valid_vars.append(v)
else:
valid_vars.append({"name": str(v), "type": "independent", "measurement_method": "TBD"})
def _ensure_str(val, default="TBD"):
if val is None: return default
if isinstance(val, list):
return "\n- ".join([str(x) for x in val]) if val else default
return str(val)
return ExperimentalProtocol(
title=_ensure_str(proto_data.get("title"), "New Protocol"),
objective=_ensure_str(proto_data.get("objective"), "Objective"),
steps=steps,
hypothesis=_ensure_str(proto_data.get("hypothesis"), "TBD"),
expected_outcomes=_ensure_str(proto_data.get("expected_outcomes"), "TBD"),
variables=valid_vars,
statistical_analysis_plan=_ensure_str(proto_data.get("statistical_analysis_plan"), "Standard descriptive statistics"),
success_criteria=proto_data.get("success_criteria", ["Completion of all steps"]),
estimated_duration_days=float(proto_data.get("estimated_duration_days", 30.0)),
alternative_approaches=proto_data.get("alternative_approaches", ["None specified"]),
risk_assessment=proto_data.get("risk_assessment", {"overall_risk": "low"})
)
def _log_reasoning(self, phase: str, step: str, description: str):
self.reasoning_trace.append({
"phase": phase, "step": step, "description": description, "timestamp": datetime.now().isoformat()
})
logger.debug(f"[{phase}] {step}: {description}")
async def chat(self, message: str, analysis_context: Optional[Dict[str, Any]] = None, history: List[Dict[str, str]] = [], user_id: Optional[str] = None) -> Dict[str, Any]:
# Minimal chat implementation for K2
llm = ChatOpenAI(model="MBZUAI-IFM/K2-Think-v2", openai_api_key=settings.K2_THINK_API_KEY, openai_api_base=settings.K2_THINK_API_URL)
resp = await llm.ainvoke([HumanMessage(content=message)])
return {"answer": resp.content, "reasoning_log": "", "suggested_actions": []}