from __future__ import annotations import json import re from typing import Any from huggingface_hub import InferenceClient from config import DEFAULT_MAX_TOKENS, HF_MODEL_ID, HF_PROVIDER, HF_TOKEN, MAX_AGENT_STEPS from data_engine import DataContext, baseline_pyspark_pipeline, baseline_sql_pipeline # System prompt fed to the model on every run. sets up the agent's role, instructs it to use tools before answering, and locks the final answer # into a fixed set of markdown sections (including fenced SQL/PySpark blocks) so downstream parsing (_extract_code) can reliably pull the code back out SYSTEM_PROMPT = '''You are a senior Data Engineering Agent. Use tools to verify dataset claims before generating pipelines. Identify concrete schema and data-quality risks, propose production-oriented fixes, and generate SQL and PySpark grounded in the uploaded dataset. Never claim generated PySpark was executed in this app. Final answer must contain exactly these sections: ## Engineering assessment ## Data-quality findings ## Recommended fixes ## SQL Pipeline ```sql ... ``` ## PySpark Pipeline ```python ... ``` SQL should use the logical table name dataset unless the user asks otherwise.''' # Tool/function definitions exposed to the model (OpenAI-style function-calling schema) so it can inspect the dataset, run checks, and validate/execute SQL before writing its final answer TOOL_SCHEMAS = [ {"type":"function","function":{"name":"inspect_dataset","description":"Inspect dataset shape, inferred types, nulls, cardinality, and examples.","parameters":{"type":"object","properties":{}}}}, {"type":"function","function":{"name":"run_quality_checks","description":"Run deterministic data-quality checks.","parameters":{"type":"object","properties":{}}}}, {"type":"function","function":{"name":"validate_expected_schema","description":"Validate the uploaded dataset against the expected JSON schema.","parameters":{"type":"object","properties":{}}}}, {"type":"function","function":{"name":"validate_sql","description":"Parse/normalize SQL with SQLGlot.","parameters":{"type":"object","properties":{"sql":{"type":"string"}},"required":["sql"]}}}, {"type":"function","function":{"name":"execute_readonly_sql","description":"Execute safe read-only DuckDB SQL against table dataset.","parameters":{"type":"object","properties":{"sql":{"type":"string"}},"required":["sql"]}}}, ] def _extract_code(text: str, language: str) -> str: # Pull the contents of a fenced code block (```sql ... ``` or ```python ... ```) out of the model's final answer. Accepts a couple of common aliases (py/pyspark) when looking for python code aliases=[language] + (["py","pyspark"] if language=="python" else []) match=re.search(rf"```(?:{'|'.join(re.escape(x) for x in aliases)})\s*(.*?)```", text or "", re.I|re.S) return match.group(1).strip() if match else "" def _serialize_tool_call(call: Any) -> dict[str, Any]: # Normalize a tool-call object returned by the HF client into a plain dict suitable for re-sending back to the model as conversation history # Arguments may arrive as a JSON string; make sure they're valid JSON (falling back to an empty dict) before re-serializing args=call.function.arguments if isinstance(args,str): try: args=json.loads(args) except json.JSONDecodeError: args={} return {"id":getattr(call,"id",None) or f"call_{call.function.name}","type":"function", "function":{"name":call.function.name,"arguments":json.dumps(args or {})}} class DataEngineeringAgent: def __init__(self, context: DataContext, expected_schema_json: str="", dialect: str="duckdb", target: str="SQL + PySpark", temperature: float=.15, max_tokens: int=DEFAULT_MAX_TOKENS): # context: wraps the uploaded dataset and exposes profiling/validation/ execution helpers used by the tool handlers below self.context=context; self.expected_schema_json=expected_schema_json or ""; self.dialect=dialect or "duckdb"; self.target=target or "SQL + PySpark" self.temperature=float(temperature); self.max_tokens=int(max_tokens); self.trace=[] # trace: log of every tool call + result for this run def _tool(self, name: str, arguments: dict[str, Any]) -> dict[str, Any]: # Dispatch a single tool call by name to the matching DataContext method, record it in self.trace, and return the result to send back to the model if name=="inspect_dataset": result=self.context.profile() elif name=="run_quality_checks": result=self.context.quality_report() elif name=="validate_expected_schema": result=self.context.validate_schema(self.expected_schema_json) elif name=="validate_sql": result=self.context.validate_sql(arguments.get("sql",""),dialect=self.dialect) elif name=="execute_readonly_sql": result=self.context.execute_sql(arguments.get("sql","")) else: result={"error":f"Unknown tool: {name}"} self.trace.append({"tool":name,"arguments":arguments,"result":result}); return result def _fallback(self, task: str, error: str|None=None): # Deterministic, non-LLM fallback path. Used when no HF_TOKEN is configured, or when the HF model call raises an exception # Runs the core tools directly and stitches together a fixed-format answer plus baseline SQL/PySpark code, so the app still produces a usable result without the model p=self._tool("inspect_dataset",{}); q=self._tool("run_quality_checks",{}); s=self._tool("validate_expected_schema",{}) sql=baseline_sql_pipeline(self.context); pyspark=baseline_pyspark_pipeline(self.context) issues="\n".join(f"- **{i['severity'].upper()}** {i.get('column') or 'dataset'}: {i['evidence']} Fix: {i['recommended_fix']}" for i in q.get("issues",[])[:6]) or "- No built-in quality issues detected." note="HF model generation is disabled because HF_TOKEN is not configured." if not HF_TOKEN else f"The HF model call failed, so deterministic fallback output was returned. Error: {error}" text=f'''## Engineering assessment {note} Task: {task} Loaded {p['source']} with **{p['rows']:,} rows** and **{p['columns']:,} columns**. ## Data-quality findings {issues} ## Recommended fixes - Enforce the schema contract before downstream writes. - Quarantine failed casts and invalid required fields. - Deduplicate with an explicit business key and deterministic ordering rule. - Add row-count, null-rate, uniqueness, and freshness checks. - Schema validator: **{s.get('message','No schema result')}** ## SQL Pipeline ```sql {sql} ``` ## PySpark Pipeline ```python {pyspark} ``` ''' return text,sql,pyspark,self.trace def run(self, task: str): # Main entry point. If no HF token is configured, skip straight to the deterministic fallback # Otherwise, drive a tool-calling loop against the HF inference client: let the model call tools (up to MAX_AGENT_STEPS rounds), feed results back as "tool" messages, and stop once it returns a plain text (non-tool-call) final answer if not HF_TOKEN: return self._fallback(task) kwargs={"token":HF_TOKEN,"timeout":120} if HF_PROVIDER: kwargs["provider"]=HF_PROVIDER client=InferenceClient(**kwargs) messages=[{"role":"system","content":SYSTEM_PROMPT},{"role":"user","content":f"Task: {task}\nTarget output: {self.target}\nSQL dialect: {self.dialect}\nExpected schema supplied: {'yes' if self.expected_schema_json.strip() else 'no'}\nUse tools first."}] try: final="" for _ in range(MAX_AGENT_STEPS): response=client.chat_completion(model=HF_MODEL_ID,messages=messages,tools=TOOL_SCHEMAS,tool_choice="auto",temperature=self.temperature,max_tokens=self.max_tokens) msg=response.choices[0].message; calls=getattr(msg,"tool_calls",None) or [] if not calls: final=(msg.content or "").strip(); break # model gave a final answer, stop looping # Record the assistant's tool-call turn, then execute each requested tool and append its result as a "tool" role message messages.append({"role":"assistant","content":msg.content or "","tool_calls":[_serialize_tool_call(c) for c in calls]}) for c in calls: args=c.function.arguments if isinstance(args,str): try: args=json.loads(args) except json.JSONDecodeError: args={} result=self._tool(c.function.name,args or {}) messages.append({"role":"tool","tool_call_id":getattr(c,"id",None) or f"call_{c.function.name}","name":c.function.name,"content":json.dumps(result,default=str)[:30000]}) if not final: # Ran out of steps without a plain-text answer; force one more call telling the model to stop calling tools and synthesize response=client.chat_completion(model=HF_MODEL_ID,messages=messages+[{"role":"user","content":"Synthesize the final answer now. Do not call more tools."}],temperature=self.temperature,max_tokens=self.max_tokens) final=(response.choices[0].message.content or "").strip() # Pull SQL/PySpark code blocks out of the model's answer; if either is missing, fall back to the deterministic baseline pipeline sql=_extract_code(final,"sql") or baseline_sql_pipeline(self.context) pyspark=_extract_code(final,"python") or baseline_pyspark_pipeline(self.context) return final,sql,pyspark,self.trace except Exception as exc: # Any failure in the model loop (network, API, parsing, etc.) falls back to the deterministic path instead of crashing return self._fallback(task,str(exc))