Files
govoplan-dataflow/tests/test_graph_and_sql.py

126 lines
4.2 KiB
Python

from __future__ import annotations
import unittest
from govoplan_dataflow.backend.executor import execute_preview
from govoplan_dataflow.backend.schemas import (
GraphEdge,
GraphNode,
GraphPosition,
PipelineGraph,
)
from govoplan_dataflow.backend.sql_compiler import SqlCompilationError, compile_sql, render_sql
def inline_source() -> GraphNode:
return GraphNode(
id="source",
type="source.inline",
label="Monthly input",
position=GraphPosition(x=40, y=160),
config={
"source_name": "monthly_files",
"rows": [
{"department": "A", "status": "open", "amount": 10},
{"department": "A", "status": "closed", "amount": 5},
{"department": "B", "status": "open", "amount": 20},
],
},
)
class DataflowGraphAndSqlTests(unittest.TestCase):
def test_compiles_renders_and_executes_grouped_query(self) -> None:
sql = """
SELECT department, COUNT(*) AS records, SUM(amount) AS total
FROM monthly_files
WHERE status = 'open'
GROUP BY department
ORDER BY total DESC
LIMIT 50
"""
graph, normalized, diagnostics = compile_sql(sql, source_nodes=[inline_source()])
self.assertEqual([], diagnostics)
self.assertEqual(
["source.inline", "filter", "aggregate", "sort", "limit", "output"],
[node.type for node in graph.nodes],
)
self.assertEqual(normalized, render_sql(graph)[0])
result = execute_preview(graph, row_limit=100)
self.assertEqual(
[
{"department": "B", "records": 1, "total": 20},
{"department": "A", "records": 1, "total": 10},
],
result.rows,
)
self.assertEqual(2, result.total_rows)
self.assertFalse(result.truncated)
self.assertEqual(3, result.input_row_count)
self.assertEqual(6, len(result.node_diagnostics))
self.assertEqual(1, len(result.source_fingerprints))
def test_rejects_effectful_or_multi_statement_sql(self) -> None:
for sql in (
"DELETE FROM monthly_files",
"SELECT * FROM monthly_files; DROP TABLE monthly_files",
"INSERT INTO target SELECT * FROM monthly_files",
):
with self.subTest(sql=sql), self.assertRaises(SqlCompilationError):
compile_sql(sql, source_nodes=[inline_source()])
def test_rejects_join_until_comparison_slice(self) -> None:
with self.assertRaisesRegex(SqlCompilationError, "JOIN support"):
compile_sql(
"SELECT * FROM monthly_files JOIN other ON monthly_files.id = other.id",
source_nodes=[inline_source()],
)
def test_graph_validation_rejects_cycle(self) -> None:
source = inline_source()
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=300, y=160),
config={},
)
graph = PipelineGraph(
nodes=[source, output],
edges=[
GraphEdge(id="forward", source="source", target="output"),
GraphEdge(id="back", source="output", target="source"),
],
)
from govoplan_dataflow.backend.graph import validate_graph
codes = {item.code for item in validate_graph(graph)}
self.assertIn("graph.cycle", codes)
def test_preview_truncates_response_without_changing_pipeline_limit(self) -> None:
source = inline_source()
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=300, y=160),
config={},
)
graph = PipelineGraph(
nodes=[source, output],
edges=[GraphEdge(id="edge", source="source", target="output")],
)
result = execute_preview(graph, row_limit=2)
self.assertEqual(2, len(result.rows))
self.assertEqual(3, result.total_rows)
self.assertTrue(result.truncated)
if __name__ == "__main__":
unittest.main()