AaTekle
Deploy Data Engineering Agent
f07ef1d
Raw History Blame Contribute Delete
9.78 kB
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))