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))