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.
Files changed (38) hide show
  1. google/cloud/db_context_enrichment/__init__.py +1 -0
  2. google/cloud/db_context_enrichment/bootstrap/__init__.py +1 -0
  3. google/cloud/db_context_enrichment/bootstrap/bootstrap_generator.py +74 -0
  4. google/cloud/db_context_enrichment/common/__init__.py +0 -0
  5. google/cloud/db_context_enrichment/common/config.py +10 -0
  6. google/cloud/db_context_enrichment/common/context_mutator.py +132 -0
  7. google/cloud/db_context_enrichment/common/parameterizer.py +192 -0
  8. google/cloud/db_context_enrichment/dataset/__init__.py +0 -0
  9. google/cloud/db_context_enrichment/dataset/dataset_generator.py +44 -0
  10. google/cloud/db_context_enrichment/evaluate/__init__.py +3 -0
  11. google/cloud/db_context_enrichment/evaluate/db_generators/__init__.py +0 -0
  12. google/cloud/db_context_enrichment/evaluate/db_generators/alloydb.py +72 -0
  13. google/cloud/db_context_enrichment/evaluate/db_generators/base.py +76 -0
  14. google/cloud/db_context_enrichment/evaluate/db_generators/mysql.py +69 -0
  15. google/cloud/db_context_enrichment/evaluate/db_generators/postgres.py +69 -0
  16. google/cloud/db_context_enrichment/evaluate/db_generators/spanner.py +62 -0
  17. google/cloud/db_context_enrichment/evaluate/evaluate_generator.py +239 -0
  18. google/cloud/db_context_enrichment/evaluate/result_reader.py +181 -0
  19. google/cloud/db_context_enrichment/facet/__init__.py +0 -0
  20. google/cloud/db_context_enrichment/facet/facet_generator.py +68 -0
  21. google/cloud/db_context_enrichment/main.py +450 -0
  22. google/cloud/db_context_enrichment/model/__init__.py +0 -0
  23. google/cloud/db_context_enrichment/model/context.py +86 -0
  24. google/cloud/db_context_enrichment/prompts/__init__.py +9 -0
  25. google/cloud/db_context_enrichment/prompts/targeted_facets.py +56 -0
  26. google/cloud/db_context_enrichment/prompts/targeted_templates.py +57 -0
  27. google/cloud/db_context_enrichment/prompts/targeted_value_search.py +93 -0
  28. google/cloud/db_context_enrichment/template/__init__.py +0 -0
  29. google/cloud/db_context_enrichment/template/template_generator.py +67 -0
  30. google/cloud/db_context_enrichment/value_search/__init__.py +0 -0
  31. google/cloud/db_context_enrichment/value_search/generator.py +91 -0
  32. google/cloud/db_context_enrichment/value_search/match_templates.py +267 -0
  33. google_cloud_db_context_engineering-0.5.1.dist-info/METADATA +160 -0
  34. google_cloud_db_context_engineering-0.5.1.dist-info/RECORD +38 -0
  35. google_cloud_db_context_engineering-0.5.1.dist-info/WHEEL +5 -0
  36. google_cloud_db_context_engineering-0.5.1.dist-info/entry_points.txt +2 -0
  37. google_cloud_db_context_engineering-0.5.1.dist-info/licenses/LICENSE +202 -0
  38. 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,10 @@
1
+ MODEL_NAME = "gemini-3.1-flash-lite"
2
+
3
+
4
+ def get_model_name() -> str:
5
+ """Returns the configured model name.
6
+
7
+ Centralized in this file to make it easy to update.
8
+ Requires rebuilding the binary to take effect.
9
+ """
10
+ return MODEL_NAME
@@ -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)}"
@@ -0,0 +1,3 @@
1
+ """
2
+ Evaluation workflow integration and configuration generation for Evalbench.
3
+ """
@@ -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()