google-cloud-db-context-engineering 0.5.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- google/cloud/db_context_enrichment/__init__.py +1 -0
- google/cloud/db_context_enrichment/bootstrap/__init__.py +1 -0
- google/cloud/db_context_enrichment/bootstrap/bootstrap_generator.py +74 -0
- google/cloud/db_context_enrichment/common/__init__.py +0 -0
- google/cloud/db_context_enrichment/common/config.py +10 -0
- google/cloud/db_context_enrichment/common/context_mutator.py +132 -0
- google/cloud/db_context_enrichment/common/parameterizer.py +192 -0
- google/cloud/db_context_enrichment/dataset/__init__.py +0 -0
- google/cloud/db_context_enrichment/dataset/dataset_generator.py +44 -0
- google/cloud/db_context_enrichment/evaluate/__init__.py +3 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/__init__.py +0 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/alloydb.py +72 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/base.py +76 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/mysql.py +69 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/postgres.py +69 -0
- google/cloud/db_context_enrichment/evaluate/db_generators/spanner.py +62 -0
- google/cloud/db_context_enrichment/evaluate/evaluate_generator.py +239 -0
- google/cloud/db_context_enrichment/evaluate/result_reader.py +181 -0
- google/cloud/db_context_enrichment/facet/__init__.py +0 -0
- google/cloud/db_context_enrichment/facet/facet_generator.py +68 -0
- google/cloud/db_context_enrichment/main.py +450 -0
- google/cloud/db_context_enrichment/model/__init__.py +0 -0
- google/cloud/db_context_enrichment/model/context.py +86 -0
- google/cloud/db_context_enrichment/prompts/__init__.py +9 -0
- google/cloud/db_context_enrichment/prompts/targeted_facets.py +56 -0
- google/cloud/db_context_enrichment/prompts/targeted_templates.py +57 -0
- google/cloud/db_context_enrichment/prompts/targeted_value_search.py +93 -0
- google/cloud/db_context_enrichment/template/__init__.py +0 -0
- google/cloud/db_context_enrichment/template/template_generator.py +67 -0
- google/cloud/db_context_enrichment/value_search/__init__.py +0 -0
- google/cloud/db_context_enrichment/value_search/generator.py +91 -0
- google/cloud/db_context_enrichment/value_search/match_templates.py +267 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/METADATA +160 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/RECORD +38 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/WHEEL +5 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/entry_points.txt +2 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/licenses/LICENSE +202 -0
- google_cloud_db_context_engineering-0.5.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# db_context_enrichment package
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# Initialization script for bootstrap module
|
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
import json
|
|
2
|
+
|
|
3
|
+
from pydantic import ValidationError
|
|
4
|
+
|
|
5
|
+
from google.cloud.db_context_enrichment.common.context_mutator import (
|
|
6
|
+
Mutation,
|
|
7
|
+
mutate_context_set,
|
|
8
|
+
)
|
|
9
|
+
from google.cloud.db_context_enrichment.facet import facet_generator
|
|
10
|
+
from google.cloud.db_context_enrichment.model import context
|
|
11
|
+
from google.cloud.db_context_enrichment.template import template_generator
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
async def generate_context(
|
|
15
|
+
output_file_path: str,
|
|
16
|
+
sql_dialect: str,
|
|
17
|
+
template_inputs_json: str | None = None,
|
|
18
|
+
facet_inputs_json: str | None = None,
|
|
19
|
+
) -> str:
|
|
20
|
+
"""
|
|
21
|
+
Core logic for generating a single unified ContextSet from key information and saving it to a file.
|
|
22
|
+
"""
|
|
23
|
+
final_templates = None
|
|
24
|
+
final_facets = None
|
|
25
|
+
|
|
26
|
+
if template_inputs_json:
|
|
27
|
+
res_str = await template_generator.generate_templates(
|
|
28
|
+
template_inputs_json, sql_dialect
|
|
29
|
+
)
|
|
30
|
+
if '"error":' in res_str:
|
|
31
|
+
raise RuntimeError(f"Error generating templates: {res_str}")
|
|
32
|
+
try:
|
|
33
|
+
res_dict = json.loads(res_str)
|
|
34
|
+
final_templates = [
|
|
35
|
+
context.Template(**t) for t in res_dict.get("templates", [])
|
|
36
|
+
]
|
|
37
|
+
except (json.JSONDecodeError, ValidationError) as e:
|
|
38
|
+
raise ValueError(f"Error parsing generated templates: {e}") from e
|
|
39
|
+
|
|
40
|
+
if facet_inputs_json:
|
|
41
|
+
res_str = await facet_generator.generate_facets(facet_inputs_json, sql_dialect)
|
|
42
|
+
if '"error":' in res_str:
|
|
43
|
+
raise RuntimeError(f"Error generating facets: {res_str}")
|
|
44
|
+
try:
|
|
45
|
+
res_dict = json.loads(res_str)
|
|
46
|
+
final_facets = [context.Facet(**f) for f in res_dict.get("facets", [])]
|
|
47
|
+
except (json.JSONDecodeError, ValidationError) as e:
|
|
48
|
+
raise ValueError(f"Error parsing generated facets: {e}") from e
|
|
49
|
+
|
|
50
|
+
mutations: list[Mutation] = []
|
|
51
|
+
|
|
52
|
+
if final_templates:
|
|
53
|
+
for t in final_templates:
|
|
54
|
+
mutations.append(
|
|
55
|
+
Mutation(
|
|
56
|
+
operation="add",
|
|
57
|
+
type="template",
|
|
58
|
+
value=t.model_dump(exclude_none=True),
|
|
59
|
+
)
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
if final_facets:
|
|
63
|
+
for f in final_facets:
|
|
64
|
+
mutations.append(
|
|
65
|
+
Mutation(
|
|
66
|
+
operation="add", type="facet", value=f.model_dump(exclude_none=True)
|
|
67
|
+
)
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
if not mutations:
|
|
71
|
+
raise ValueError("No templates or facets were generated to save.")
|
|
72
|
+
|
|
73
|
+
mutate_context_set(output_file_path, mutations)
|
|
74
|
+
return output_file_path
|
|
File without changes
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
from typing import Any, Literal
|
|
4
|
+
|
|
5
|
+
from pydantic import BaseModel, ValidationError
|
|
6
|
+
|
|
7
|
+
from google.cloud.db_context_enrichment.model import context
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Mutation(BaseModel):
|
|
11
|
+
"""Defines a strict schema for context modifications."""
|
|
12
|
+
|
|
13
|
+
operation: Literal["add", "delete", "update"]
|
|
14
|
+
type: Literal["template", "facet", "value_search"]
|
|
15
|
+
identifier: dict[str, Any] = {}
|
|
16
|
+
value: dict[str, Any] | None = None
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def mutate_context_set(file_path: str, mutations: list[Mutation]) -> None:
|
|
20
|
+
"""
|
|
21
|
+
Internal function to mutate (add, delete, update) elements in an existing ContextSet JSON file.
|
|
22
|
+
|
|
23
|
+
The mutations is expected to be a list of Mutation objects.
|
|
24
|
+
Each object specifies an 'operation', 'type', 'identifier' (to route to the correct item for deletes/updates), and 'value' (for adding/updating), for instance:
|
|
25
|
+
|
|
26
|
+
[
|
|
27
|
+
{
|
|
28
|
+
"operation": "add",
|
|
29
|
+
"type": "template",
|
|
30
|
+
"value": {"nl_query": "...", "sql": "...", "intent": "...", "manifest": "...", "parameterized": {...}}
|
|
31
|
+
},
|
|
32
|
+
{
|
|
33
|
+
"operation": "delete",
|
|
34
|
+
"type": "template",
|
|
35
|
+
"identifier": {"nl_query": "What are all users?"}
|
|
36
|
+
},
|
|
37
|
+
{
|
|
38
|
+
"operation": "update",
|
|
39
|
+
"type": "facet",
|
|
40
|
+
"identifier": {"intent": "high price"},
|
|
41
|
+
"value": {"sql_snippet": "price > 2000", "intent": "very high price"}
|
|
42
|
+
}
|
|
43
|
+
]
|
|
44
|
+
"""
|
|
45
|
+
|
|
46
|
+
# 1. Load exiting ContextSet (or create an empty one)
|
|
47
|
+
if not os.path.exists(file_path) or os.path.getsize(file_path) == 0:
|
|
48
|
+
context_set = context.ContextSet()
|
|
49
|
+
else:
|
|
50
|
+
try:
|
|
51
|
+
with open(file_path) as f:
|
|
52
|
+
raw_data = json.load(f)
|
|
53
|
+
context_set = context.ContextSet.model_validate(raw_data)
|
|
54
|
+
except ValidationError as e:
|
|
55
|
+
raise ValueError(
|
|
56
|
+
f"Validation Error loading ContextSet from {file_path}: {e}"
|
|
57
|
+
) from e
|
|
58
|
+
except (json.JSONDecodeError, OSError) as e:
|
|
59
|
+
raise RuntimeError(f"Error reading JSON from {file_path}: {e}") from e
|
|
60
|
+
|
|
61
|
+
# 2. Model mapping for tracking and validation
|
|
62
|
+
type_to_model = {
|
|
63
|
+
"template": context.Template,
|
|
64
|
+
"facet": context.Facet,
|
|
65
|
+
"value_search": context.ValueSearch,
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
type_to_attr = {
|
|
69
|
+
"template": "templates",
|
|
70
|
+
"facet": "facets",
|
|
71
|
+
"value_search": "value_searches",
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
# 3. Apply mutations
|
|
75
|
+
for i, mut in enumerate(mutations):
|
|
76
|
+
op = mut.operation
|
|
77
|
+
item_type = mut.type
|
|
78
|
+
identifier = mut.identifier
|
|
79
|
+
value_data = mut.value
|
|
80
|
+
|
|
81
|
+
if item_type not in type_to_attr:
|
|
82
|
+
continue
|
|
83
|
+
|
|
84
|
+
attr_name = type_to_attr[item_type]
|
|
85
|
+
model_class = type_to_model[item_type]
|
|
86
|
+
|
|
87
|
+
target_list = getattr(context_set, attr_name)
|
|
88
|
+
if target_list is None:
|
|
89
|
+
target_list = []
|
|
90
|
+
setattr(context_set, attr_name, target_list)
|
|
91
|
+
|
|
92
|
+
if op == "add":
|
|
93
|
+
if value_data:
|
|
94
|
+
try:
|
|
95
|
+
new_item = model_class.model_validate(value_data)
|
|
96
|
+
target_list.append(new_item)
|
|
97
|
+
except ValidationError as e:
|
|
98
|
+
raise ValueError(
|
|
99
|
+
f"Validation Error on mutation {i} during 'add': {e}"
|
|
100
|
+
) from e
|
|
101
|
+
|
|
102
|
+
elif op == "delete":
|
|
103
|
+
new_list = []
|
|
104
|
+
for item in target_list:
|
|
105
|
+
item_dict = item.model_dump()
|
|
106
|
+
match = all(item_dict.get(k) == v for k, v in identifier.items())
|
|
107
|
+
if not match:
|
|
108
|
+
new_list.append(item)
|
|
109
|
+
setattr(context_set, attr_name, new_list)
|
|
110
|
+
|
|
111
|
+
elif op == "update":
|
|
112
|
+
for idx, item in enumerate(target_list):
|
|
113
|
+
item_dict = item.model_dump()
|
|
114
|
+
match = all(item_dict.get(k) == v for k, v in identifier.items())
|
|
115
|
+
if match and value_data:
|
|
116
|
+
updated_dict = {**item_dict, **value_data}
|
|
117
|
+
try:
|
|
118
|
+
updated_item = model_class.model_validate(updated_dict)
|
|
119
|
+
target_list[idx] = updated_item
|
|
120
|
+
break # Only update first match
|
|
121
|
+
except ValidationError as e:
|
|
122
|
+
raise ValueError(
|
|
123
|
+
f"Validation Error on mutation {i} during 'update': {e}"
|
|
124
|
+
) from e
|
|
125
|
+
|
|
126
|
+
# 4. Save validated ContextSet
|
|
127
|
+
try:
|
|
128
|
+
os.makedirs(os.path.dirname(os.path.abspath(file_path)), exist_ok=True)
|
|
129
|
+
with open(file_path, "w") as f:
|
|
130
|
+
f.write(context_set.model_dump_json(indent=2, exclude_none=True))
|
|
131
|
+
except OSError as e:
|
|
132
|
+
raise RuntimeError(f"Error saving ContextSet to {file_path}: {e}") from e
|
|
@@ -0,0 +1,192 @@
|
|
|
1
|
+
import re
|
|
2
|
+
import textwrap
|
|
3
|
+
from enum import Enum
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from pydantic import BaseModel, Field
|
|
7
|
+
|
|
8
|
+
from google import genai
|
|
9
|
+
from google.cloud.db_context_enrichment.common import config
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class SQLDialect(Enum):
|
|
13
|
+
"""Enumeration for supported database dialects."""
|
|
14
|
+
|
|
15
|
+
POSTGRESQL = "postgresql"
|
|
16
|
+
MYSQL = "mysql"
|
|
17
|
+
GOOGLESQL = "googlesql"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class ValuePhrasePair(BaseModel):
|
|
21
|
+
"""A key-value pair for a named entity and its types."""
|
|
22
|
+
|
|
23
|
+
key: str = Field(..., description="The extracted named entity.")
|
|
24
|
+
value: list[str] = Field(
|
|
25
|
+
..., description="A list of identified types for the entity."
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class ValuePhrasesList(BaseModel):
|
|
30
|
+
"""A list of named entity pairs."""
|
|
31
|
+
|
|
32
|
+
value_phrases: list[ValuePhrasePair] = Field(
|
|
33
|
+
...,
|
|
34
|
+
description="A list of key-value pairs, where each key is a named entity and the value is a list of its types.",
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
async def extract_value_phrases(nl_query: str) -> dict[str, list[str]]:
|
|
39
|
+
"""
|
|
40
|
+
Extracts potential value phrases from a natural language question using an LLM.
|
|
41
|
+
|
|
42
|
+
This function replicates the core logic of the `value_phrases_extractor`
|
|
43
|
+
and `get_value_phrases_extractor_template` functions in `choose.sql`.
|
|
44
|
+
It builds a prompt to perform named entity recognition (NER) and calls a
|
|
45
|
+
generative model to extract entities based on a predefined list of types.
|
|
46
|
+
|
|
47
|
+
Args:
|
|
48
|
+
nl_question: The natural language question to analyze.
|
|
49
|
+
|
|
50
|
+
Returns:
|
|
51
|
+
A dictionary containing the extracted phrases and their types.
|
|
52
|
+
|
|
53
|
+
Raises:
|
|
54
|
+
Exception: If the model call fails or returns an invalid response.
|
|
55
|
+
"""
|
|
56
|
+
prompt = textwrap.dedent(
|
|
57
|
+
f"""
|
|
58
|
+
Please extract the named entity (a real-world object, such as a person,
|
|
59
|
+
location, organization, product, etc., that can be denoted with a proper name)
|
|
60
|
+
from the query literally based on the following types:
|
|
61
|
+
|
|
62
|
+
[Types]
|
|
63
|
+
- country
|
|
64
|
+
- city
|
|
65
|
+
- email_address
|
|
66
|
+
- language
|
|
67
|
+
- law
|
|
68
|
+
- organization
|
|
69
|
+
- person
|
|
70
|
+
- product
|
|
71
|
+
- sport or activity
|
|
72
|
+
- work of art
|
|
73
|
+
- date
|
|
74
|
+
- time
|
|
75
|
+
- number
|
|
76
|
+
- currency
|
|
77
|
+
- region
|
|
78
|
+
|
|
79
|
+
The output should be a JSON object containing a list of key-value pairs.
|
|
80
|
+
Each pair should have a "key" (the extracted named entity) and a "value" (a list of its identified types).
|
|
81
|
+
For example: {{"value_phrases": [{{"key": "entity1", "value": ["type1"]}}, {{"key": "entity2", "value": ["type2", "type3"]}}]}}
|
|
82
|
+
DO NOT perform any spell checking or correction. If no entities are
|
|
83
|
+
identified, return an empty list.
|
|
84
|
+
|
|
85
|
+
[Query]
|
|
86
|
+
{nl_query}
|
|
87
|
+
"""
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
client = genai.Client()
|
|
91
|
+
try:
|
|
92
|
+
response = await client.aio.models.generate_content(
|
|
93
|
+
model=config.get_model_name(),
|
|
94
|
+
contents=prompt,
|
|
95
|
+
config={
|
|
96
|
+
"response_mime_type": "application/json",
|
|
97
|
+
"response_schema": ValuePhrasesList,
|
|
98
|
+
},
|
|
99
|
+
)
|
|
100
|
+
if response.text:
|
|
101
|
+
phrases_obj = ValuePhrasesList.model_validate_json(response.text)
|
|
102
|
+
# Convert the list of pairs back to a dictionary
|
|
103
|
+
return {pair.key: pair.value for pair in phrases_obj.value_phrases}
|
|
104
|
+
else:
|
|
105
|
+
# Return an empty dict if the model returns no text
|
|
106
|
+
return {}
|
|
107
|
+
except Exception as e:
|
|
108
|
+
# Re-raise the exception to be handled by the caller
|
|
109
|
+
raise Exception(f"An error occurred during value phrase extraction: {e}") from e
|
|
110
|
+
finally:
|
|
111
|
+
client.close()
|
|
112
|
+
await client.aio.aclose()
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def parameterize_sql_and_intent(
|
|
116
|
+
value_phrases: dict[str, Any],
|
|
117
|
+
sql: str,
|
|
118
|
+
intent: str,
|
|
119
|
+
db_dialect: SQLDialect = SQLDialect.POSTGRESQL,
|
|
120
|
+
) -> dict[str, str]:
|
|
121
|
+
"""
|
|
122
|
+
Replaces value phrases in a SQL query and an intent string with placeholders.
|
|
123
|
+
|
|
124
|
+
This function iterates through a dictionary of value phrases and replaces
|
|
125
|
+
their occurrences in both the SQL and intent strings with positional
|
|
126
|
+
parameters. The syntax of the parameters is determined by the specified
|
|
127
|
+
database dialect.
|
|
128
|
+
|
|
129
|
+
The phrases are processed in descending order of length to handle nested
|
|
130
|
+
phrases correctly (e.g., "New York" before "York").
|
|
131
|
+
|
|
132
|
+
The replacement logic handles both quoted and unquoted occurrences of the
|
|
133
|
+
phrases, and it avoids replacing phrases that are already part of a
|
|
134
|
+
placeholder.
|
|
135
|
+
|
|
136
|
+
Args:
|
|
137
|
+
value_phrases: A dictionary where keys are the string phrases to be
|
|
138
|
+
replaced. Values are not used.
|
|
139
|
+
sql: The SQL query string to parameterize.
|
|
140
|
+
intent: The natural language intent string to parameterize.
|
|
141
|
+
db_dialect: The SQL dialect to use for parameterization.
|
|
142
|
+
|
|
143
|
+
Returns:
|
|
144
|
+
A dictionary containing the parameterized 'sql' and 'intent' strings.
|
|
145
|
+
"""
|
|
146
|
+
psql = sql
|
|
147
|
+
pintent = intent
|
|
148
|
+
param_index = 1
|
|
149
|
+
|
|
150
|
+
# Sort keys by length in descending order to prioritize longer matches
|
|
151
|
+
sorted_phrases = sorted(value_phrases.keys(), key=len, reverse=True)
|
|
152
|
+
|
|
153
|
+
for value in sorted_phrases:
|
|
154
|
+
# Determine the placeholder based on the database dialect
|
|
155
|
+
if db_dialect == SQLDialect.POSTGRESQL:
|
|
156
|
+
placeholder = f"${param_index}"
|
|
157
|
+
else: # For mysql, googlesql, etc.
|
|
158
|
+
placeholder = "?"
|
|
159
|
+
|
|
160
|
+
search_phrase_quoted = f"'{value}'"
|
|
161
|
+
replaced = False
|
|
162
|
+
|
|
163
|
+
# Patterns with negative lookbehind to avoid replacing existing params
|
|
164
|
+
# e.g., don't replace 'foo' in "$'foo'"
|
|
165
|
+
quoted_pattern = re.compile(r"(?<!\$)" + re.escape(search_phrase_quoted))
|
|
166
|
+
unquoted_pattern = re.compile(r"(?<!\$)" + r"\b" + re.escape(value) + r"\b")
|
|
167
|
+
|
|
168
|
+
# Condition 1: Quoted in SQL, Quoted in Intent
|
|
169
|
+
if quoted_pattern.search(psql) and quoted_pattern.search(pintent):
|
|
170
|
+
psql = quoted_pattern.sub(placeholder, psql)
|
|
171
|
+
pintent = quoted_pattern.sub(placeholder, pintent)
|
|
172
|
+
replaced = True
|
|
173
|
+
# Condition 2: Quoted in SQL, Unquoted in Intent
|
|
174
|
+
elif quoted_pattern.search(psql) and unquoted_pattern.search(pintent):
|
|
175
|
+
psql = quoted_pattern.sub(placeholder, psql)
|
|
176
|
+
pintent = unquoted_pattern.sub(placeholder, pintent)
|
|
177
|
+
replaced = True
|
|
178
|
+
# Condition 3: Unquoted in SQL, Quoted in Intent
|
|
179
|
+
elif unquoted_pattern.search(psql) and quoted_pattern.search(pintent):
|
|
180
|
+
psql = unquoted_pattern.sub(placeholder, psql)
|
|
181
|
+
pintent = quoted_pattern.sub(placeholder, pintent)
|
|
182
|
+
replaced = True
|
|
183
|
+
# Condition 4: Unquoted in SQL, Unquoted in Intent
|
|
184
|
+
elif unquoted_pattern.search(psql) and unquoted_pattern.search(pintent):
|
|
185
|
+
psql = unquoted_pattern.sub(placeholder, psql)
|
|
186
|
+
pintent = unquoted_pattern.sub(placeholder, pintent)
|
|
187
|
+
replaced = True
|
|
188
|
+
|
|
189
|
+
if replaced:
|
|
190
|
+
param_index += 1
|
|
191
|
+
|
|
192
|
+
return {"sql": psql, "intent": pintent}
|
|
File without changes
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
async def generate_dataset(
|
|
6
|
+
dataset_entries_json: str,
|
|
7
|
+
output_file_path: str,
|
|
8
|
+
) -> str:
|
|
9
|
+
"""
|
|
10
|
+
Validates a list of evaluation dataset entries and saves them to a JSON file.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
dataset_entries_json: A JSON string representing a list of dataset items.
|
|
14
|
+
Each item should have "id", "database", "nlq", and "golden_sql" keys.
|
|
15
|
+
Example: '[{"id": "eval_001", "database": "my_db", "nlq": "Count users", "golden_sql": "SELECT COUNT(*) FROM users"}]'
|
|
16
|
+
output_file_path: The absolute path where the dataset JSON file should be saved.
|
|
17
|
+
|
|
18
|
+
Returns:
|
|
19
|
+
The absolute file path where the dataset was saved.
|
|
20
|
+
"""
|
|
21
|
+
try:
|
|
22
|
+
data = json.loads(dataset_entries_json)
|
|
23
|
+
if not isinstance(data, list):
|
|
24
|
+
raise ValueError("Dataset entries must be a list of objects.")
|
|
25
|
+
|
|
26
|
+
# Simple validation of keys
|
|
27
|
+
for i, entry in enumerate(data):
|
|
28
|
+
if not isinstance(entry, dict):
|
|
29
|
+
raise ValueError(f"Entry at index {i} is not an object.")
|
|
30
|
+
missing_keys = {"id", "database", "nlq", "golden_sql"} - set(entry.keys())
|
|
31
|
+
if missing_keys:
|
|
32
|
+
raise ValueError(
|
|
33
|
+
f"Entry at index {i} is missing required keys: {missing_keys}"
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
# Ensure directory exists
|
|
37
|
+
os.makedirs(os.path.dirname(os.path.abspath(output_file_path)), exist_ok=True)
|
|
38
|
+
|
|
39
|
+
with open(output_file_path, "w") as f:
|
|
40
|
+
json.dump(data, f, indent=2)
|
|
41
|
+
|
|
42
|
+
return f"Successfully saved dataset to {output_file_path}"
|
|
43
|
+
except (json.JSONDecodeError, ValueError, OSError) as e:
|
|
44
|
+
return f"Error saving dataset: {str(e)}"
|
|
File without changes
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
from typing import Any
|
|
2
|
+
|
|
3
|
+
import google.cloud.geminidataanalytics_v1beta as gda
|
|
4
|
+
import yaml
|
|
5
|
+
|
|
6
|
+
from .base import BaseDBConfigGenerator
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class AlloyDBConfigGenerator(BaseDBConfigGenerator):
|
|
10
|
+
"""
|
|
11
|
+
Dedicated generator mapping properties to explicit AlloyDB configuration
|
|
12
|
+
topologies utilized by both EvalBench binaries and GDA Context objects.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
SOURCE_TYPE = "alloydb-postgres"
|
|
16
|
+
DIALECT = "postgres"
|
|
17
|
+
REQUIRED_FIELDS = BaseDBConfigGenerator.REQUIRED_FIELDS | {
|
|
18
|
+
"project",
|
|
19
|
+
"region",
|
|
20
|
+
"cluster",
|
|
21
|
+
"instance",
|
|
22
|
+
"database",
|
|
23
|
+
}
|
|
24
|
+
|
|
25
|
+
def __init__(self, params: dict[str, Any]):
|
|
26
|
+
super().__init__(params)
|
|
27
|
+
self.project = params.get("project")
|
|
28
|
+
self.region = params.get("region")
|
|
29
|
+
self.cluster = params.get("cluster")
|
|
30
|
+
self.instance = params.get("instance")
|
|
31
|
+
self.database = params.get("database")
|
|
32
|
+
self.user = params.get("user")
|
|
33
|
+
self.password = params.get("password")
|
|
34
|
+
|
|
35
|
+
def generate_db_config(self) -> str:
|
|
36
|
+
db_type = "alloydb"
|
|
37
|
+
db_path = f"projects/{self.project}/locations/{self.region}/clusters/{self.cluster}/instances/{self.instance}"
|
|
38
|
+
|
|
39
|
+
db_config = {
|
|
40
|
+
"db_type": db_type,
|
|
41
|
+
"dialect": self.DIALECT,
|
|
42
|
+
"database_name": self.database,
|
|
43
|
+
"database_path": db_path,
|
|
44
|
+
"max_executions_per_minute": 180,
|
|
45
|
+
"nl_config": "", # Required by evalbench schema
|
|
46
|
+
}
|
|
47
|
+
if self.user:
|
|
48
|
+
db_config["user_name"] = self.user
|
|
49
|
+
if self.password:
|
|
50
|
+
db_config["password"] = self.password
|
|
51
|
+
return yaml.safe_dump(
|
|
52
|
+
db_config, sort_keys=False, default_flow_style=False
|
|
53
|
+
).strip()
|
|
54
|
+
|
|
55
|
+
def build_datasource_reference(
|
|
56
|
+
self, context_set_id: str
|
|
57
|
+
) -> gda.DatasourceReferences:
|
|
58
|
+
datasource_ref = gda.DatasourceReferences()
|
|
59
|
+
|
|
60
|
+
datasource_ref.alloydb = gda.AlloyDbReference(
|
|
61
|
+
database_reference=gda.AlloyDbDatabaseReference(
|
|
62
|
+
project_id=self.project,
|
|
63
|
+
region=self.region,
|
|
64
|
+
cluster_id=self.cluster,
|
|
65
|
+
instance_id=self.instance,
|
|
66
|
+
database_id=self.database,
|
|
67
|
+
),
|
|
68
|
+
agent_context_reference=gda.AgentContextReference(
|
|
69
|
+
context_set_id=context_set_id
|
|
70
|
+
),
|
|
71
|
+
)
|
|
72
|
+
return datasource_ref
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import Any
|
|
3
|
+
|
|
4
|
+
import google.cloud.geminidataanalytics_v1beta as gda
|
|
5
|
+
import yaml
|
|
6
|
+
from google.protobuf.json_format import MessageToDict
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class BaseDBConfigGenerator(ABC):
|
|
10
|
+
"""
|
|
11
|
+
Abstract Base Class enforcing the construction contract for Evalbench database topologies.
|
|
12
|
+
Each distinct DB type (Spanner, Postgres, AlloyDB, MySQL) must inherit and implement
|
|
13
|
+
the mappings required by both the standard EvalBench framework and the GDA SDK model.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
SOURCE_TYPE = "unknown"
|
|
17
|
+
DIALECT = "unknown"
|
|
18
|
+
REQUIRED_FIELDS = {"project"}
|
|
19
|
+
|
|
20
|
+
def __init__(self, params: dict[str, Any]):
|
|
21
|
+
self.params = params
|
|
22
|
+
self.validate()
|
|
23
|
+
|
|
24
|
+
@abstractmethod
|
|
25
|
+
def generate_db_config(self) -> str:
|
|
26
|
+
"""
|
|
27
|
+
Generates the Evalbench db_config.yaml payload natively.
|
|
28
|
+
"""
|
|
29
|
+
raise NotImplementedError("Subclasses must implement generate_db_config")
|
|
30
|
+
|
|
31
|
+
@abstractmethod
|
|
32
|
+
def build_datasource_reference(
|
|
33
|
+
self, context_set_id: str
|
|
34
|
+
) -> gda.DatasourceReferences:
|
|
35
|
+
"""
|
|
36
|
+
Constructs the strict Protocol Buffer DatasourceReference required by the QueryDataAPI
|
|
37
|
+
context generation flow.
|
|
38
|
+
"""
|
|
39
|
+
raise NotImplementedError(
|
|
40
|
+
"Subclasses must implement build_datasource_reference"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
def validate(self) -> None:
|
|
44
|
+
"""
|
|
45
|
+
Validates that the provided tools.yaml source configuration block contains all
|
|
46
|
+
the mandatory fields required by the specific Evalbench topology.
|
|
47
|
+
"""
|
|
48
|
+
missing = [f for f in self.REQUIRED_FIELDS if f not in self.params]
|
|
49
|
+
if missing:
|
|
50
|
+
raise ValueError(
|
|
51
|
+
f"Missing required fields in tools.yaml config for '{self.SOURCE_TYPE}': "
|
|
52
|
+
f"{', '.join(missing)}"
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
def generate_model_config(self, context_set_id: str) -> str:
|
|
56
|
+
"""
|
|
57
|
+
Standardized Model Builder converting the strictly typed GDA object into an EvalBench model dict.
|
|
58
|
+
"""
|
|
59
|
+
datasource_ref = self.build_datasource_reference(context_set_id)
|
|
60
|
+
|
|
61
|
+
query_context = gda.QueryDataContext(datasource_references=datasource_ref)
|
|
62
|
+
|
|
63
|
+
query_context_dict = MessageToDict(
|
|
64
|
+
query_context._pb, preserving_proto_field_name=True
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
model_config = {
|
|
68
|
+
"generator": "query_data_api",
|
|
69
|
+
"project_id": self.params.get("project"),
|
|
70
|
+
"location": self.params.get("region") or "global",
|
|
71
|
+
"context": query_context_dict,
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
return yaml.safe_dump(
|
|
75
|
+
model_config, sort_keys=False, default_flow_style=False
|
|
76
|
+
).strip()
|