|
15 | 15 | # limitations under the License. |
16 | 16 | # |
17 | 17 | """Unit tests for generative model prompts.""" |
| 18 | + |
18 | 19 | # pylint: disable=protected-access,bad-continuation |
19 | 20 |
|
| 21 | +from google.cloud.aiplatform import initializer as aiplatform_initializer |
| 22 | +from google.cloud.aiplatform.compat.types import dataset as gca_dataset |
| 23 | +from google.cloud.aiplatform_v1.types import ( |
| 24 | + dataset_version as gca_dataset_version, |
| 25 | +) |
20 | 26 | from vertexai.prompts._prompts import Prompt |
| 27 | +from vertexai.prompts import _prompt_management |
21 | 28 | from vertexai.generative_models import ( |
22 | 29 | Content, |
23 | 30 | Part, |
|
42 | 49 | types_v1 as gapic_tool_types, |
43 | 50 | ) |
44 | 51 |
|
45 | | - |
46 | 52 | _RESPONSE_TEXT_PART_STRUCT = { |
47 | 53 | "text": "The sky appears blue due to a phenomenon called Rayleigh scattering." |
48 | 54 | } |
@@ -305,6 +311,75 @@ def create_image(): |
305 | 311 | class TestPrompt: |
306 | 312 | """Unit tests for generative model prompts.""" |
307 | 313 |
|
| 314 | + def _make_prompt_dataset_metadata(self): |
| 315 | + prompt = Prompt(prompt_data="Rate the movie {movie}", model_name="gemini-pro") |
| 316 | + return _prompt_management._format_dataset_metadata_dict(prompt=prompt) |
| 317 | + |
| 318 | + @mock.patch.object(Prompt, "_dataset_client", new_callable=mock.PropertyMock) |
| 319 | + def test_get_latest_prompt_populates_version_metadata(self, dataset_client_mock): |
| 320 | + aiplatform_initializer.global_config.init( |
| 321 | + project="test-project", location="us-central1" |
| 322 | + ) |
| 323 | + prompt_id = "123456789" |
| 324 | + dataset_name = "projects/test-project/locations/us-central1/datasets/123456789" |
| 325 | + metadata = self._make_prompt_dataset_metadata() |
| 326 | + dataset_client = mock.Mock() |
| 327 | + dataset_client.get_dataset.return_value = gca_dataset.Dataset( |
| 328 | + name=dataset_name, |
| 329 | + display_name="test prompt", |
| 330 | + metadata_schema_uri=_prompt_management.PROMPT_SCHEMA_URI, |
| 331 | + metadata=metadata, |
| 332 | + model_reference="gemini-pro", |
| 333 | + ) |
| 334 | + dataset_client.list_dataset_versions.return_value = [ |
| 335 | + gca_dataset_version.DatasetVersion( |
| 336 | + name=f"{dataset_name}/datasetVersions/3", |
| 337 | + display_name="version 3", |
| 338 | + ) |
| 339 | + ] |
| 340 | + dataset_client_mock.return_value = dataset_client |
| 341 | + |
| 342 | + prompt = _prompt_management.get(prompt_id=prompt_id) |
| 343 | + |
| 344 | + assert prompt.version_id == "3" |
| 345 | + assert prompt.version_name == "version 3" |
| 346 | + dataset_client.list_dataset_versions.assert_called_once_with( |
| 347 | + parent=dataset_name, |
| 348 | + page_size=1, |
| 349 | + order_by="create_time desc", |
| 350 | + ) |
| 351 | + |
| 352 | + @mock.patch.object(Prompt, "_dataset_client", new_callable=mock.PropertyMock) |
| 353 | + def test_get_pinned_prompt_does_not_list_latest_version(self, dataset_client_mock): |
| 354 | + aiplatform_initializer.global_config.init( |
| 355 | + project="test-project", location="us-central1" |
| 356 | + ) |
| 357 | + prompt_id = "123456789" |
| 358 | + version_id = "2" |
| 359 | + dataset_name = "projects/test-project/locations/us-central1/datasets/123456789" |
| 360 | + version_name = f"{dataset_name}/datasetVersions/{version_id}" |
| 361 | + metadata = self._make_prompt_dataset_metadata() |
| 362 | + dataset_client = mock.Mock() |
| 363 | + dataset_client.get_dataset_version.return_value = ( |
| 364 | + gca_dataset_version.DatasetVersion( |
| 365 | + name=version_name, |
| 366 | + display_name="version 2", |
| 367 | + metadata=metadata, |
| 368 | + model_reference="gemini-pro", |
| 369 | + ) |
| 370 | + ) |
| 371 | + dataset_client.get_dataset.return_value = gca_dataset.Dataset( |
| 372 | + name=dataset_name, |
| 373 | + display_name="test prompt", |
| 374 | + ) |
| 375 | + dataset_client_mock.return_value = dataset_client |
| 376 | + |
| 377 | + prompt = _prompt_management.get(prompt_id=prompt_id, version_id=version_id) |
| 378 | + |
| 379 | + assert prompt.version_id == version_id |
| 380 | + assert prompt.version_name == "version 2" |
| 381 | + dataset_client.list_dataset_versions.assert_not_called() |
| 382 | + |
308 | 383 | def test_string_prompt_constructor_string_variables(self): |
309 | 384 | # Create string prompt with string only variable values |
310 | 385 | prompt = Prompt( |
|
0 commit comments