Spaces:
Running on Zero
Running on Zero
Download agent.py from AaronTekle/Agentic-DataEngineering: direct link, hf CLI and curl.
- Browser
- Download file 9.78 kB
-
https://huggingface.co/spaces/AaronTekle/Agentic-DataEngineering/resolve/main/agent.py
- Command line
-
hf download hf://spaces/AaronTekle/Agentic-DataEngineering/agent.py
-
curl -L -o agent.py https://huggingface.co/spaces/AaronTekle/Agentic-DataEngineering/resolve/main/agent.py
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)) |