Spaces:
Running on Zero
Running on Zero
File size: 9,784 Bytes
f07ef1d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | 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)) |