Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 68 additions & 32 deletions actions/generate_storyboard.py
Original file line number Diff line number Diff line change
Expand Up @@ -285,19 +285,25 @@ def execute(

# Group images by product ID to build context
products: dict[str, Product] = {}
valid_product_image_combinations = {}

for img_obj in images:
product_id = str(img_obj.get(Dimension.PRODUCT_ID.value, "1"))
product_description = str(img_obj.get("product_description", None))
desc = img_obj.get("product_description")
product_description = str(desc) if desc is not None else ""
if product_id not in products:
products[product_id] = Product(product_id, product_description)
if product_id not in valid_product_image_combinations:
valid_product_image_combinations[product_id] = set()

image_id = str(img_obj.get(Dimension.IMAGE_ID.value, "1"))
product = products[product_id]
image = Image(
id=str(img_obj.get(Dimension.IMAGE_ID.value, "1")),
id=image_id,
uri=gcs.get_uri(img_obj[Key.FILE.value]),
)
product.images.append(image)
valid_product_image_combinations[product_id].add(image_id)

# Add products and images
prompt_parts.append(genai.types.Part.from_text(text="### Products & Images:\n\n"))
Expand Down Expand Up @@ -373,41 +379,71 @@ def execute(
],
)

start_time = time.time()
response = client.models.generate_content(
model=gemini_model,
contents=prompt_parts,
config=config,
)
end_time = time.time()
logger.info(
"Gemini API request completed in %.2f seconds.", end_time - start_time
)

segments = []
if (
response.candidates
and response.candidates[0].content
and response.candidates[0].content.parts
):
parts = list(response.candidates[0].content.parts)
for part in parts:
if part.text:
segments.append(part.text)
storyboard_text = "".join(segments)
if not storyboard_text:
raise ValueError(
"The model returned no storyboard text. This can happen when the"
" model reaches the output token limit before emitting a response."
result = None
for _ in range(4): # 1 initial call + 3 retries
start_time = time.time()
response = client.models.generate_content(
model=gemini_model,
contents=prompt_parts,
config=config,
)
end_time = time.time()
logger.info(
"Gemini API request completed in %.2f seconds.", end_time - start_time
)
json_result = json.loads(storyboard_text)

if not json_result.get("storyboard"):
raise ValueError("No storyboard found in the response.")
segments = []
if (
response.candidates
and response.candidates[0].content
and response.candidates[0].content.parts
):
parts = list(response.candidates[0].content.parts)
for part in parts:
if part.text:
segments.append(part.text)
storyboard_text = "".join(segments)
if not storyboard_text:
raise ValueError(
"The model returned no storyboard text. This can happen when the"
" model reaches the output token limit before emitting a response."
)
candidate_result = json.loads(storyboard_text)

if not candidate_result.get("storyboard"):
raise ValueError("No storyboard found in the response.")

is_valid = True
for scene in candidate_result["storyboard"]:
product_id = scene.get(Dimension.PRODUCT_ID.value)
image_id = scene.get(Dimension.IMAGE_ID.value)
video_prompt = scene.get("video_prompt")
scene_name = scene.get("scene_name")
if (
product_id not in valid_product_image_combinations
or image_id
not in valid_product_image_combinations.get(product_id, set())
or not video_prompt
or not scene_name
):
is_valid = False
logger.warning(
"Invalid scene from Gemini: product_id=%s, image_id=%s",
product_id,
image_id,
)
break # This scene is invalid, so the whole result is.
Comment on lines +417 to +435

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In response_schema, video_prompt and scene_name are typed as "string" without "minLength": 1. Currently, the retry loop only checks (product_id, image_id) against valid_product_image_combinations and immediately breaks out of the retry loop with is_valid = True even if video_prompt or scene_name is an empty string (""). Those scenes are then silently dropped after the loop at lines 447–448 (if not image_id or not product_id or not video_prompt or not scene_name: continue), which can leave valid_scenes empty ({"storyboard": []}) without having triggered a retry.

Including not video_prompt and not scene_name in the retry validation condition ensures empty scene fields trigger the retry loop instead of bypassing retries and being silently discarded post-loop.

Suggested change
for scene in candidate_result["storyboard"]:
product_id = scene.get(Dimension.PRODUCT_ID.value)
image_id = scene.get(Dimension.IMAGE_ID.value)
if (
product_id not in valid_product_image_combinations
or image_id
not in valid_product_image_combinations.get(product_id, set())
):
is_valid = False
logger.warning(
"Invalid scene from Gemini: product_id=%s, image_id=%s",
product_id,
image_id,
)
break # This scene is invalid, so the whole result is.
for scene in candidate_result["storyboard"]:
product_id = scene.get(Dimension.PRODUCT_ID.value)
image_id = scene.get(Dimension.IMAGE_ID.value)
video_prompt = scene.get("video_prompt")
scene_name = scene.get("scene_name")
if (
product_id not in valid_product_image_combinations
or image_id
not in valid_product_image_combinations.get(product_id, set())
or not video_prompt
or not scene_name
):
is_valid = False
logger.warning(
"Invalid scene from Gemini: product_id=%s, image_id=%s",
product_id,
image_id,
)
break # This scene is invalid, so the whole result is.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed, applied in a273ab6. An empty video_prompt or scene_name now fails the in-loop validation and triggers a retry, instead of passing validation and then being dropped by the post-loop filter (which could store {"storyboard": []} without ever retrying).

New test test_execute_empty_scene_fields_trigger_retry covers both fields as subtests. It asserts a second Gemini call and that the stored storyboard keeps its scene. Both subtests fail against the previous commit and pass now. Full Python suite: 713 passed.

if is_valid:
result = candidate_result
break
if not result:
raise RuntimeError(
"Gemini repeatedly scripted invalid image/product combinations"
)

valid_scenes = []

for scene in json_result["storyboard"]:
for scene in result["storyboard"]:
image_id = scene.get(Dimension.IMAGE_ID.value)
product_id = scene.get(Dimension.PRODUCT_ID.value)
video_prompt = scene.get("video_prompt")
Expand Down
229 changes: 229 additions & 0 deletions actions/test/test_generate_storyboard.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,235 @@

class TestGenerateStoryboard(unittest.TestCase):

def _make_mock_response(self, storyboard_dict):
mock_response = MagicMock()
mock_candidate = MagicMock()
mock_part = MagicMock()
mock_part.text = json.dumps(storyboard_dict)
mock_candidate.content.parts = [mock_part]
mock_response.candidates = [mock_candidate]
return mock_response

@patch("actions.generate_storyboard.genai.Client")
def test_execute_hallucinated_product_id_triggers_retry(
self, mock_genai_client_class
):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
invalid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i1",
"product_id": "p_invalid",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
valid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i1",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
mock_client.models.generate_content.side_effect = [invalid_resp, valid_resp]

result = generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
self.mock_images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

self.assertEqual(mock_client.models.generate_content.call_count, 2)
self.assertEqual(
result,
{"storyboard": [{Key.FILE.value: "gs://bucket/storyboard.json"}]},
)

@patch("actions.generate_storyboard.genai.Client")
def test_execute_hallucinated_image_id_triggers_retry(
self, mock_genai_client_class
):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
invalid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i_invalid",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
valid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i1",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
mock_client.models.generate_content.side_effect = [invalid_resp, valid_resp]

result = generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
self.mock_images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

self.assertEqual(mock_client.models.generate_content.call_count, 2)
self.assertEqual(
result,
{"storyboard": [{Key.FILE.value: "gs://bucket/storyboard.json"}]},
)

@patch("actions.generate_storyboard.genai.Client")
def test_execute_empty_scene_fields_trigger_retry(
self, mock_genai_client_class
):
for field in ("video_prompt", "scene_name"):
with self.subTest(field=field):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
scene = {
"image_id": "i1",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}
empty_resp = self._make_mock_response(
{"storyboard": [{**scene, field: ""}]}
)
valid_resp = self._make_mock_response({"storyboard": [scene]})
mock_client.models.generate_content.side_effect = [
empty_resp,
valid_resp,
]
self.mock_gcs.store.reset_mock()

generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
self.mock_images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

self.assertEqual(mock_client.models.generate_content.call_count, 2)
stored = json.loads(self.mock_gcs.store.call_args.args[0])
self.assertEqual(len(stored["storyboard"]), 1)

@patch("actions.generate_storyboard.genai.Client")
def test_execute_exhausting_retries_raises_runtime_error(
self, mock_genai_client_class
):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
invalid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i_invalid",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
mock_client.models.generate_content.return_value = invalid_resp

with self.assertRaises(RuntimeError) as cm:
generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
self.mock_images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

self.assertEqual(mock_client.models.generate_content.call_count, 4)
self.assertIn("Gemini repeatedly scripted invalid image/product combinations", str(cm.exception))

@patch("actions.generate_storyboard.genai.Client")
def test_execute_valid_ids_pass_through_unchanged(
self, mock_genai_client_class
):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
valid_resp = self._make_mock_response({
"storyboard": [{
"image_id": "i1",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})
mock_client.models.generate_content.return_value = valid_resp

generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
self.mock_images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

self.mock_gcs.store.assert_called_once()
stored_json = json.loads(self.mock_gcs.store.call_args[0][0])
self.assertEqual(stored_json["storyboard"][0]["product_id"], "p1")
self.assertEqual(stored_json["storyboard"][0]["image_id"], "i1")

@patch("actions.generate_storyboard.genai.Client")
def test_execute_missing_product_description_yields_empty_string_not_none(
self, mock_genai_client_class
):
mock_client = MagicMock()
mock_genai_client_class.return_value = mock_client
mock_client.models.generate_content.return_value = self._make_mock_response({
"storyboard": [{
"image_id": "i1",
"product_id": "p1",
"scene_name": "Scene 1",
"video_prompt": "Prompt 1",
}]
})

# Test with both omitted and explicitly None product_description
images = [
{
Dimension.PRODUCT_ID.value: "p1",
Dimension.IMAGE_ID.value: "i1",
Key.FILE.value: "path/to/img1.jpg",
# product_description omitted
},
{
Dimension.PRODUCT_ID.value: "p2",
Dimension.IMAGE_ID.value: "i2",
Key.FILE.value: "path/to/img2.jpg",
"product_description": None,
},
]

generate_storyboard.execute(
self.mock_gcs,
self.mock_params,
images,
self.mock_user_prompt,
self.gemini_model,
self.gemini_model_location,
)

_, kwargs = mock_client.models.generate_content.call_args
prompt_parts = kwargs["contents"]
all_text = "".join(p.text for p in prompt_parts if getattr(p, "text", None))
self.assertNotIn("'None'", all_text)
self.assertNotIn("Product description:", all_text)

def setUp(self):
self.mock_gcs = MagicMock()
self.mock_params = {Key.GCP_PROJECT.value: "test-project"}
Expand Down
Loading
Loading