import os import json import re from groq import Groq class NagaMLOpsAgent: def __init__( self, api_key: str = "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", model: str = "llama-3.3-70b-versatile" ): self.api_key = api_key self.primary_model = model self.fallback_models = [ "llama-3.3-70b-versatile", "llama-3.1-8b-instant", "deepseek-r1-distill-llama-70b", "mixtral-8x7b-32768" ] self.client = Groq(api_key=self.api_key) def diagnose_and_heal( self, script_code: str, execution_logs: str, telemetry: dict, fault_name: str = None ) -> dict: """ Sends error logs, broken python code, and telemetry context to Groq ultra-fast API. Returns a structured dictionary containing root cause diagnosis and executable patched code. """ system_prompt = ( "You are an expert Autonomous MLOps & AI Infrastructure Diagnostic Agent.\n" "Your task is to analyze failing machine learning pipelines, identify root causes, " "and generate production-grade, fully working Python code patches to fix the issue.\n\n" "CRITICAL INSTRUCTION: You MUST return your answer in valid JSON format matching this EXACT schema:\n" "{\n" ' "fault_category": "DATA_DRIFT | CODE_RUNTIME_ERROR | NAN_LOSS | OOM_SPIKE | MODEL_ACCURACY_DROP",\n' ' "severity": "CRITICAL | HIGH | MEDIUM",\n' ' "root_cause_analysis": "Detailed explanation of why the crash/degradation happened.",\n' ' "explanation_for_engineers": "Actionable summary for MLOps dashboard.",\n' ' "patch_code": "FULL valid Python script replacing the broken code completely without placeholder comments.",\n' ' "verification_checklist": ["Check 1", "Check 2"]\n' "}" ) user_content = ( f"=== REPORTED FAULT SCENARIO ===\n{fault_name or 'Auto-Detected Anomaly'}\n\n" f"=== PIPELINE TELEMETRY ===\n{json.dumps(telemetry, indent=2)}\n\n" f"=== BROKEN PIPELINE CODE ===\n```python\n{script_code}\n```\n\n" f"=== EXECUTION LOGS & STACKTRACE ===\n{execution_logs}\n\n" "Perform root cause analysis and produce the fixed `patch_code` Python script. " "Ensure the patch is self-contained, syntax-correct, and completely resolves the error." ) response_text = "" last_error = None models_to_try = [self.primary_model] + [m for m in self.fallback_models if m != self.primary_model] for m in models_to_try: try: completion = self.client.chat.completions.create( model=m, messages=[ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_content} ], temperature=0.1, response_format={"type": "json_object"} ) response_text = completion.choices[0].message.content.strip() if response_text: print(f"[AgentBrain] Super-fast Groq diagnosis generated using model: {m}") break except Exception as e: print(f"[AgentBrain] Model {m} failed: {e}. Trying next fallback...") last_error = e if not response_text: return self._generate_fallback_diagnosis(script_code, str(last_error)) return self._parse_agent_response(response_text, script_code) def _parse_agent_response(self, raw_text: str, original_code: str) -> dict: clean_text = raw_text if "```json" in clean_text: clean_text = clean_text.split("```json")[1].split("```")[0].strip() elif "```" in clean_text and clean_text.strip().startswith("```"): clean_text = clean_text.split("```")[1].split("```")[0].strip() try: parsed = json.loads(clean_text) patch = parsed.get("patch_code", original_code) if isinstance(patch, dict): patch = patch.get("code", patch.get("script", str(patch))) elif not isinstance(patch, str): patch = str(patch) if "```python" in patch: patch = patch.split("```python")[1].split("```")[0].strip() elif "```" in patch: patch = patch.split("```")[1].split("```")[0].strip() parsed["patch_code"] = str(patch) return parsed except Exception as json_err: print(f"[AgentBrain] JSON parsing failed: {json_err}. Extracting code via regex fallback...") code_match = re.search(r"```python(.*?)```", raw_text, re.DOTALL) patch_code = code_match.group(1).strip() if code_match else original_code return { "fault_category": "CODE_RUNTIME_ERROR", "severity": "HIGH", "root_cause_analysis": raw_text[:300] + "...", "explanation_for_engineers": "Agent generated fix. Extracted code patch successfully.", "patch_code": patch_code, "verification_checklist": ["Execute patched script", "Verify pipeline telemetry"] } def _generate_fallback_diagnosis(self, original_code: str, error_msg: str) -> dict: return { "fault_category": "CODE_RUNTIME_ERROR", "severity": "HIGH", "root_cause_analysis": f"Local Fallback Diagnosis: Pipeline exception detected ({error_msg}).", "explanation_for_engineers": "Network issue reaching LLM API. Initiated local safety patch.", "patch_code": original_code.replace("/ 0", "/ 1.0").replace("np.nan", "0.0"), "verification_checklist": ["Local syntax check", "Re-run safety sandbox"] }