Skip to content

Commit 9dd474c

Browse files
committed
fix(prompts): populate latest version metadata
1 parent f6ecd67 commit 9dd474c

2 files changed

Lines changed: 95 additions & 2 deletions

File tree

tests/unit/vertexai/test_prompts.py

Lines changed: 76 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,9 +15,16 @@
1515
# limitations under the License.
1616
#
1717
"""Unit tests for generative model prompts."""
18+
1819
# pylint: disable=protected-access,bad-continuation
1920

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+
)
2026
from vertexai.prompts._prompts import Prompt
27+
from vertexai.prompts import _prompt_management
2128
from vertexai.generative_models import (
2229
Content,
2330
Part,
@@ -42,7 +49,6 @@
4249
types_v1 as gapic_tool_types,
4350
)
4451

45-
4652
_RESPONSE_TEXT_PART_STRUCT = {
4753
"text": "The sky appears blue due to a phenomenon called Rayleigh scattering."
4854
}
@@ -305,6 +311,75 @@ def create_image():
305311
class TestPrompt:
306312
"""Unit tests for generative model prompts."""
307313

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+
308383
def test_string_prompt_constructor_string_variables(self):
309384
# Create string prompt with string only variable values
310385
prompt = Prompt(

vertexai/prompts/_prompt_management.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -500,6 +500,22 @@ def _get_prompt_resource(prompt: Prompt, prompt_id: str) -> gca_dataset.Dataset:
500500
return dataset
501501

502502

503+
def _populate_latest_version_metadata(prompt: Prompt, prompt_id: str) -> None:
504+
"""Populates prompt version metadata for the latest prompt version."""
505+
project = aiplatform_initializer.global_config.project
506+
location = aiplatform_initializer.global_config.location
507+
parent = f"projects/{project}/locations/{location}/datasets/{prompt_id}"
508+
versions = prompt._dataset_client.list_dataset_versions(
509+
parent=parent,
510+
page_size=1,
511+
order_by="create_time desc",
512+
)
513+
for version in versions:
514+
prompt._version_id = version.name.split("/")[-1]
515+
prompt._version_name = version.display_name
516+
break
517+
518+
503519
def _get_prompt_resource_from_version(
504520
prompt: Prompt, prompt_id: str, version_id: str
505521
) -> gca_dataset.Dataset:
@@ -580,12 +596,14 @@ def get(prompt_id: str, version_id: Optional[str] = None) -> Prompt:
580596
)
581597
else:
582598
dataset = _get_prompt_resource(prompt=prompt, prompt_id=prompt_id)
599+
_populate_latest_version_metadata(prompt=prompt, prompt_id=prompt_id)
583600

584601
# Remove etag to avoid error for repeated dataset updates
585602
dataset.etag = None
586603

587604
prompt._dataset = dataset
588-
prompt._version_id = version_id
605+
if version_id:
606+
prompt._version_id = version_id
589607

590608
dataset_dict = _proto_to_dict(dataset)
591609

0 commit comments

Comments
 (0)