diff --git a/src/ecommerce_agent/dashboard/design_lineage.py b/src/ecommerce_agent/dashboard/design_lineage.py new file mode 100644 index 0000000..e20b29c --- /dev/null +++ b/src/ecommerce_agent/dashboard/design_lineage.py @@ -0,0 +1,45 @@ +from typing import Any + + +def children_by_parent(generations: list[dict[str, Any]]) -> dict[str | None, list[dict[str, Any]]]: + children: dict[str | None, list[dict[str, Any]]] = {} + for generation in generations: + children.setdefault(generation.get("parent_generation_id"), []).append(generation) + return children + + +def leaf_generations(generations: list[dict[str, Any]]) -> list[dict[str, Any]]: + parent_ids = { + str(generation["parent_generation_id"]) + for generation in generations + if generation.get("parent_generation_id") + } + return [generation for generation in generations if str(generation["id"]) not in parent_ids] + + +def lineage_path(leaf: dict[str, Any], by_id: dict[str, dict[str, Any]]) -> list[dict[str, Any]]: + path = [leaf] + seen = {str(leaf["id"])} + current = leaf + while current.get("parent_generation_id") in by_id: + parent = by_id[str(current["parent_generation_id"])] + if str(parent["id"]) in seen: + break + path.append(parent) + seen.add(str(parent["id"])) + current = parent + return list(reversed(path)) + + +def lineage_paths(generations: list[dict[str, Any]]) -> list[list[dict[str, Any]]]: + by_id = {str(item["id"]): item for item in generations} + paths = [lineage_path(leaf, by_id) for leaf in leaf_generations(generations)] + return sorted(paths, key=lambda path: path[-1]["created_at"], reverse=True) + + +def version_label(path: list[dict[str, Any]], generation: dict[str, Any]) -> str: + return f"v{path.index(generation) + 1}" + + +def branch_label(path: list[dict[str, Any]]) -> str: + return " -> ".join(f"v{index}" for index in range(1, len(path) + 1)) diff --git a/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py b/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py index 7a97ee3..9bd4526 100644 --- a/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py +++ b/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py @@ -4,6 +4,7 @@ import streamlit as st from ecommerce_agent.dashboard.client import delete, get, get_asset, patch, post_multipart +from ecommerce_agent.dashboard.design_lineage import branch_label, lineage_paths, version_label from ecommerce_agent.dashboard.ui import page_header, show_error from ecommerce_agent.domain.dtos import ImageGenerationRequest, ImageReference from ecommerce_agent.services.openai_images import build_effective_prompt @@ -34,8 +35,11 @@ "design_aspect": "Square 1:1", "design_custom_width": 4, "design_custom_height": 5, + "focused_generation_id": None, + "focused_lineage_leaf_id": None, "editing_generation_id": None, "deleting_generation_id": None, + "deleting_lineage_leaf_id": None, } for key, value in DEFAULT_STATE.items(): if key not in st.session_state: @@ -75,6 +79,22 @@ def confirm_delete(generation_id: str) -> None: st.session_state.deleting_generation_id = None +def start_delete_lineage(leaf_id: str) -> None: + st.session_state.deleting_lineage_leaf_id = leaf_id + + +def cancel_delete_lineage() -> None: + st.session_state.deleting_lineage_leaf_id = None + + +def confirm_delete_lineage(path: list[dict[str, Any]]) -> None: + for generation in reversed(path): + delete(f"/image-generations/{generation['id']}") + if st.session_state.editing_generation_id in {item["id"] for item in path}: + st.session_state.editing_generation_id = None + st.session_state.deleting_lineage_leaf_id = None + + def upload_files(uploaded: list[Any]) -> list[tuple[str, tuple[str, bytes, str]]]: return [ ( @@ -107,42 +127,14 @@ def show_generated_image( return content -def lineage_groups(generations: list[dict[str, Any]]) -> list[list[dict[str, Any]]]: - by_id = {item["id"]: item for item in generations} - groups: dict[str, list[dict[str, Any]]] = {} - for generation in generations: - root = generation - seen: set[str] = set() - while root.get("parent_generation_id") in by_id and root["id"] not in seen: - seen.add(root["id"]) - root = by_id[root["parent_generation_id"]] - groups.setdefault(root["id"], []).append(generation) - return sorted( - (sorted(group, key=lambda item: item["created_at"]) for group in groups.values()), - key=lambda group: group[-1]["created_at"], - reverse=True, - ) - - -def preview_versions(group: list[dict[str, Any]]) -> list[dict[str, Any]]: - if len(group) <= 4: - return group - indexes = sorted({0, len(group) // 3, (len(group) * 2) // 3, len(group) - 1}) - return [group[index] for index in indexes] - - -def version_label(group: list[dict[str, Any]], generation: dict[str, Any]) -> str: - return f"v{group.index(generation) + 1}" +def line_title(path: list[dict[str, Any]]) -> str: + leaf = path[-1] + return str(leaf.get("title") or "Untitled design") -def line_title(group: list[dict[str, Any]]) -> str: - return next((str(item["title"]) for item in group if item.get("title")), "Untitled design") - - -def save_line_title(group: list[dict[str, Any]], title: str) -> None: +def save_line_title(path: list[dict[str, Any]], title: str) -> None: clean_title = title.strip() or "Untitled design" - for generation in group: - patch(f"/image-generations/{generation['id']}", {"title": clean_title}) + patch(f"/image-generations/{path[-1]['id']}", {"title": clean_title}) def aspect_dimensions( @@ -227,62 +219,268 @@ def _format_label(output_format: str) -> str: def show_details(generation: dict[str, Any]) -> None: - with st.expander("Details"): - st.caption( - f"{generation['id']} · {generation['model']} · {generation['quality']} · " - f"{generation['width']}x{generation['height']} · " - f"{generation['output_format'].upper()} · " - f"estimated ${generation['estimated_cost_usd']}" - ) - if generation["background_removal_friendly"]: - st.caption("Background-removal friendly") - if generation.get("usage"): - st.json(generation["usage"], expanded=False) - if generation.get("error_message"): - st.error(generation["error_message"]) - - -def generation_actions(generation: dict[str, Any]) -> None: - actions = st.columns(3) + st.caption( + f"{generation['id']} · {generation['model']} · {generation['quality']} · " + f"{generation['width']}x{generation['height']} · " + f"{generation['output_format'].upper()} · " + f"estimated ${generation['estimated_cost_usd']}" + ) + if generation["background_removal_friendly"]: + st.caption("Background-removal friendly") + if generation.get("usage"): + st.json(generation["usage"], expanded=False) + if generation.get("error_message"): + st.error(generation["error_message"]) + + +def thumbnail_rows( + path: list[dict[str, Any]], *, key_prefix: str, open_dialog: bool = True +) -> None: + for row_start in range(0, len(path), 4): + row = path[row_start : row_start + 4] + widths = [1] * len(row) + if len(row) < 4: + widths.append(4 - len(row)) + columns = st.columns(widths, gap="small") + for column, generation in zip(columns, row, strict=False): + with column: + if generation.get("asset_uri"): + show_generated_image(generation["asset_uri"], width=120) + else: + st.caption(generation["status"]) + label = version_label(path, generation) + st.caption(label) + if open_dialog: + if st.button("Open", key=f"{key_prefix}-open-{generation['id']}"): + version_dialog(path, generation) + elif st.button("Select", key=f"{key_prefix}-select-{generation['id']}"): + st.session_state.focused_generation_id = generation["id"] + + +def version_actions( + path: list[dict[str, Any]], generation: dict[str, Any], *, key_prefix: str +) -> None: + extension, content_type = OUTPUT_DOWNLOADS[generation["output_format"]] + image_content = get_asset(generation["asset_uri"]) if generation.get("asset_uri") else b"" + actions = st.columns(4) with actions[0]: - st.button( - "Reuse", - key=f"reuse-{generation['id']}", - on_click=reuse_generation, - args=(generation,), - ) - with actions[1]: - st.button( - "Edit", - key=f"edit-{generation['id']}", + if st.button( + f"Edit from {version_label(path, generation)}", + key=f"{key_prefix}-edit-{generation['id']}", disabled=generation["status"] != "succeeded" or not generation.get("asset_uri"), - on_click=start_edit, - args=(generation["id"],), - ) + ): + start_edit(generation["id"]) + st.rerun() + with actions[1]: + if st.button("Reuse prompt", key=f"{key_prefix}-reuse-{generation['id']}"): + reuse_generation(generation) + st.success("Prompt copied to New design.") with actions[2]: - st.button( - "Delete", - key=f"delete-{generation['id']}", - on_click=start_delete, - args=(generation["id"],), + st.download_button( + "Download", + data=image_content, + file_name=( + f"{line_title(path).lower().replace(' ', '-')}-" + f"{version_label(path, generation)}.{extension}" + ), + mime=content_type, + key=f"{key_prefix}-download-{generation['id']}", + disabled=not bool(image_content), + on_click="ignore", ) + with actions[3]: + if st.button("Delete", key=f"{key_prefix}-delete-{generation['id']}"): + start_delete(generation["id"]) + if st.session_state.deleting_generation_id == generation["id"]: st.warning("Delete this version? Edited descendants will be preserved.") confirm, cancel = st.columns(2) with confirm: - st.button( + if st.button( "Confirm delete", - key=f"confirm-delete-{generation['id']}", + key=f"{key_prefix}-confirm-delete-{generation['id']}", + type="primary", + ): + confirm_delete(generation["id"]) + st.rerun() + with cancel: + if st.button("Cancel", key=f"{key_prefix}-cancel-delete-{generation['id']}"): + cancel_delete() + st.rerun() + + +@st.dialog("Design version", width="large") +def version_dialog(path: list[dict[str, Any]], generation: dict[str, Any]) -> None: + title = line_title(path) + label = version_label(path, generation) + st.subheader(f"{title} · {label}") + image_col, info_col = st.columns([1, 1]) + with image_col: + if generation.get("asset_uri"): + show_generated_image(generation["asset_uri"]) + else: + st.caption(f"Generation status: {generation['status']}") + with info_col: + st.write(generation["prompt"]) + show_details(generation) + version_actions(path, generation, key_prefix=f"version-dialog-{generation['id']}") + + +@st.dialog("Design lineage", width="large") +def lineage_dialog(path: list[dict[str, Any]], lineage_summary: str) -> None: + title = line_title(path) + st.subheader(f"{title} · Lineage") + st.caption(lineage_summary) + rename, save = st.columns([3, 1]) + with rename: + new_title = st.text_input("Design title", value=title, key=f"title-{path[-1]['id']}") + with save: + st.write("") + if st.button("Save title", key=f"save-title-{path[-1]['id']}"): + save_line_title(path, new_title) + st.rerun() + + thumbnail_rows(path, key_prefix=f"lineage-dialog-{path[-1]['id']}", open_dialog=False) + + selected = next( + (item for item in path if item["id"] == st.session_state.focused_generation_id), + path[-1], + ) + st.divider() + st.caption(f"Selected: {version_label(path, selected)}") + image_col, info_col = st.columns([1, 1]) + with image_col: + if selected.get("asset_uri"): + show_generated_image(selected["asset_uri"]) + else: + st.caption(f"Generation status: {selected['status']}") + with info_col: + st.write(selected["prompt"]) + show_details(selected) + version_actions(path, selected, key_prefix=f"lineage-selected-{selected['id']}") + + +@st.dialog("Edit design", width="large") +def edit_dialog(selected: dict[str, Any], path: list[dict[str, Any]] | None) -> None: + label = version_label(path, selected) if path is not None and selected in path else "version" + st.subheader(f"Edit from {selected.get('title') or 'Untitled design'} · {label}") + edit_left, edit_right = st.columns([1, 1]) + with edit_left: + show_generated_image(selected["asset_uri"]) + st.caption(f"Source {selected['width']}x{selected['height']}") + with edit_right: + edit_instruction = st.text_area( + "Edit instruction", + key=f"edit-instruction-{selected['id']}", + height=150, + placeholder="Describe what should change while preserving the rest...", + ) + edit_references = st.file_uploader( + "Additional reference images", + type=["png", "jpg", "jpeg", "webp"], + accept_multiple_files=True, + key=f"edit-references-{selected['id']}", + help="Upload up to three additional references.", + ) + show_upload_previews(edit_references) + edit_aspect = st.radio( + "Aspect ratio", + EDIT_ASPECT_OPTIONS, + key=f"edit-aspect-{selected['id']}", + horizontal=True, + ) + if edit_aspect == "Custom": + ratio_width, ratio_height = ratio_parts(selected["width"], selected["height"]) + custom_left, custom_right = st.columns(2) + with custom_left: + edit_custom_width = st.number_input( + "Ratio width", + min_value=1, + max_value=20, + value=ratio_width, + key=f"edit-custom-width-{selected['id']}", + ) + with custom_right: + edit_custom_height = st.number_input( + "Ratio height", + min_value=1, + max_value=20, + value=ratio_height, + key=f"edit-custom-height-{selected['id']}", + ) + else: + edit_custom_width = selected["width"] + edit_custom_height = selected["height"] + edit_width, edit_height = aspect_dimensions( + edit_aspect, + edit_custom_width, + edit_custom_height, + source=selected, + ) + st.caption(f"Output size: {edit_width}x{edit_height}") + edit_format_label = st.selectbox( + "Edited output format", + options=list(OUTPUT_FORMATS), + index=list(OUTPUT_FORMATS).index(_format_label(selected["output_format"])), + key=f"edit-format-{selected['id']}", + ) + edit_background = st.checkbox( + "Background-removal friendly", + value=selected["background_removal_friendly"], + key=f"edit-background-{selected['id']}", + ) + too_many_edit = len(edit_references) > 3 + if too_many_edit: + st.error("A maximum of three additional reference images is allowed.") + preview_edit, render, cancel = st.columns(3) + with preview_edit: + show_edit_preview = st.button( + "Preview full prompt", + disabled=not edit_instruction.strip(), + key=f"preview-edit-prompt-{selected['id']}", + ) + with render: + render_edit = st.button( + f"Generate from {label}", type="primary", - on_click=confirm_delete, - args=(generation["id"],), + disabled=not edit_instruction.strip() or too_many_edit, ) with cancel: - st.button( - "Cancel", - key=f"cancel-delete-{generation['id']}", - on_click=cancel_delete, + if st.button("Cancel", on_click=cancel_edit): + st.rerun() + if show_edit_preview: + st.code( + prompt_preview( + prompt=edit_instruction.strip(), + width=edit_width, + height=edit_height, + output_format=OUTPUT_FORMATS[edit_format_label], + background_removal_friendly=edit_background, + reference_count=len(edit_references) + 1, + ) ) + if render_edit: + with st.spinner("Rendering edit with OpenAI..."): + result = post_multipart( + f"/image-generations/{selected['id']}/edits", + data={ + "prompt": edit_instruction.strip(), + "width": str(edit_width), + "height": str(edit_height), + "output_format": OUTPUT_FORMATS[edit_format_label], + "background_removal_friendly": str(edit_background).lower(), + }, + files=upload_files(edit_references), + timeout=150, + ) + if result["status"] == "succeeded": + st.success("Edited version generated and saved.") + st.session_state.editing_generation_id = None + elif result["status"] == "unknown": + st.warning(result["error_message"]) + else: + st.error(result["error_message"]) + st.rerun() try: @@ -385,205 +583,65 @@ def generation_actions(generation: dict[str, Any]) -> None: st.error(result["error_message"]) st.rerun() + paths = lineage_paths(generations) selected = next( (item for item in generations if item["id"] == st.session_state.editing_generation_id), None, ) if selected is not None: - st.divider() - st.subheader(f"Editing: {selected.get('title') or 'Untitled design'}") - edit_left, edit_right = st.columns([1, 1]) - with edit_left: - show_generated_image(selected["asset_uri"]) - st.caption(f"Source {selected['width']}x{selected['height']}") - with edit_right: - edit_instruction = st.text_area( - "Edit instruction", - key=f"edit-instruction-{selected['id']}", - height=150, - placeholder="Describe what should change while preserving the rest...", - ) - edit_references = st.file_uploader( - "Additional reference images", - type=["png", "jpg", "jpeg", "webp"], - accept_multiple_files=True, - key=f"edit-references-{selected['id']}", - help="Upload up to three additional references.", - ) - show_upload_previews(edit_references) - edit_aspect = st.radio( - "Aspect ratio", - EDIT_ASPECT_OPTIONS, - key=f"edit-aspect-{selected['id']}", - horizontal=True, - ) - if edit_aspect == "Custom": - ratio_width, ratio_height = ratio_parts(selected["width"], selected["height"]) - custom_left, custom_right = st.columns(2) - with custom_left: - edit_custom_width = st.number_input( - "Ratio width", - min_value=1, - max_value=20, - value=ratio_width, - key=f"edit-custom-width-{selected['id']}", - ) - with custom_right: - edit_custom_height = st.number_input( - "Ratio height", - min_value=1, - max_value=20, - value=ratio_height, - key=f"edit-custom-height-{selected['id']}", - ) - else: - edit_custom_width = selected["width"] - edit_custom_height = selected["height"] - edit_width, edit_height = aspect_dimensions( - edit_aspect, - edit_custom_width, - edit_custom_height, - source=selected, - ) - st.caption(f"Output size: {edit_width}x{edit_height}") - edit_format_label = st.selectbox( - "Edited output format", - options=list(OUTPUT_FORMATS), - index=list(OUTPUT_FORMATS).index(_format_label(selected["output_format"])), - key=f"edit-format-{selected['id']}", - ) - edit_background = st.checkbox( - "Background-removal friendly", - value=selected["background_removal_friendly"], - key=f"edit-background-{selected['id']}", - ) - too_many_edit = len(edit_references) > 3 - if too_many_edit: - st.error("A maximum of three additional reference images is allowed.") - preview_edit, render, cancel = st.columns(3) - with preview_edit: - show_edit_preview = st.button( - "Preview full prompt", - disabled=not edit_instruction.strip(), - key=f"preview-edit-prompt-{selected['id']}", - ) - with render: - render_edit = st.button( - "Generate next version", - type="primary", - disabled=not edit_instruction.strip() or too_many_edit, - ) - with cancel: - st.button("Cancel", on_click=cancel_edit) - if show_edit_preview: - st.code( - prompt_preview( - prompt=edit_instruction.strip(), - width=edit_width, - height=edit_height, - output_format=OUTPUT_FORMATS[edit_format_label], - background_removal_friendly=edit_background, - reference_count=len(edit_references) + 1, - ) - ) - if render_edit: - with st.spinner("Rendering edit with OpenAI..."): - result = post_multipart( - f"/image-generations/{selected['id']}/edits", - data={ - "prompt": edit_instruction.strip(), - "width": str(edit_width), - "height": str(edit_height), - "output_format": OUTPUT_FORMATS[edit_format_label], - "background_removal_friendly": str(edit_background).lower(), - }, - files=upload_files(edit_references), - timeout=150, - ) - if result["status"] == "succeeded": - st.success("Edited version generated and saved.") - st.session_state.editing_generation_id = None - elif result["status"] == "unknown": - st.warning(result["error_message"]) - else: - st.error(result["error_message"]) - st.rerun() + selected_path = next((path for path in paths if selected in path), None) + edit_dialog(selected, selected_path) st.divider() st.subheader("Design Library") if not generations: st.info("No designs have been generated.") - for group in lineage_groups(generations): - latest = group[-1] - title = line_title(group) + root_counts: dict[str, int] = {} + for path in paths: + root_counts[str(path[0]["id"])] = root_counts.get(str(path[0]["id"]), 0) + 1 + root_seen: dict[str, int] = {} + for path in paths: + latest = path[-1] + title = line_title(path) + root_id = str(path[0]["id"]) + root_seen[root_id] = root_seen.get(root_id, 0) + 1 + lineage_summary = branch_label(path) + if root_counts[root_id] > 1: + lineage_summary = ( + f"{lineage_summary} · Fork {root_seen[root_id]} of {root_counts[root_id]}" + ) with st.container(border=True): - header, actions = st.columns([3, 1]) - with header: - st.markdown(f"**{title}**") - st.caption( - f"{len(group)} version{'s' if len(group) != 1 else ''} · " - f"Latest {latest['width']}x{latest['height']} · " - f"Updated {latest['created_at']}" - ) - with actions: - st.button( - "Edit latest", - key=f"edit-latest-{latest['id']}", - disabled=latest["status"] != "succeeded" or not latest.get("asset_uri"), - on_click=start_edit, - args=(latest["id"],), - ) - - thumbnails = st.columns(len(preview_versions(group))) - for column, generation in zip(thumbnails, preview_versions(group), strict=True): - with column: - if generation.get("asset_uri"): - show_generated_image(generation["asset_uri"], width=120) - else: - st.caption(generation["status"]) - st.caption(version_label(group, generation)) - - expanded = st.checkbox("Show all versions", key=f"expand-line-{group[0]['id']}") - if expanded: - rename, save = st.columns([3, 1]) - with rename: - new_title = st.text_input( - "Design title", - value=title, - key=f"title-{group[0]['id']}", - ) - with save: - st.write("") - if st.button("Save title", key=f"save-title-{group[0]['id']}"): - save_line_title(group, new_title) + st.markdown(f"**{title}**") + st.caption( + f"{len(path)} version{'s' if len(path) != 1 else ''} · " + f"Latest {version_label(path, latest)} · Updated {latest['created_at']}" + ) + st.caption(lineage_summary) + thumbnail_rows(path, key_prefix=f"library-{latest['id']}") + inspect_col, delete_col = st.columns([1, 1]) + with inspect_col: + if st.button("View lineage", key=f"view-lineage-{latest['id']}"): + st.session_state.focused_lineage_leaf_id = latest["id"] + st.session_state.focused_generation_id = latest["id"] + lineage_dialog(path, lineage_summary) + with delete_col: + if st.session_state.deleting_lineage_leaf_id == latest["id"]: + if st.button( + "Confirm delete", + key=f"confirm-delete-lineage-{latest['id']}", + type="primary", + ): + confirm_delete_lineage(path) st.rerun() - - for generation in group: - st.divider() - image_col, info_col = st.columns([1, 2]) - with image_col: - if generation.get("asset_uri"): - image_content = show_generated_image(generation["asset_uri"]) - extension, content_type = OUTPUT_DOWNLOADS[generation["output_format"]] - st.download_button( - "Download", - data=image_content, - file_name=( - f"{title.lower().replace(' ', '-')}-" - f"{version_label(group, generation)}.{extension}" - ), - mime=content_type, - key=f"download-{generation['id']}", - on_click="ignore", - width="stretch", - ) - else: - st.caption(f"Generation status: {generation['status']}") - with info_col: - st.caption(version_label(group, generation)) - st.write(generation["prompt"]) - show_details(generation) - generation_actions(generation) + if st.button("Cancel", key=f"cancel-delete-lineage-{latest['id']}"): + cancel_delete_lineage() + st.rerun() + elif st.button( + "Delete...", + key=f"delete-lineage-{latest['id']}", + ): + start_delete_lineage(latest["id"]) + st.rerun() except Exception as error: show_error(error) diff --git a/tests/smoke/test_api_and_dashboard.py b/tests/smoke/test_api_and_dashboard.py index b99a723..69ec095 100644 --- a/tests/smoke/test_api_and_dashboard.py +++ b/tests/smoke/test_api_and_dashboard.py @@ -12,6 +12,7 @@ from ecommerce_agent.api.main import create_app from ecommerce_agent.config import Settings +from ecommerce_agent.dashboard.design_lineage import lineage_paths from ecommerce_agent.db.models import Job, ResearchRun from ecommerce_agent.db.session import get_session from ecommerce_agent.domain.enums import JobStatus, ResearchRunStatus @@ -434,6 +435,86 @@ async def _get_session(): assert result["background_removal_friendly"] is False +async def test_editing_non_latest_generation_creates_forked_lineages( + tmp_path, db_sessions, services +) -> None: + app = create_app(make_settings(tmp_path)) + app.state.services = services + + async def _get_session(): + async with db_sessions() as session: + yield session + + app.dependency_overrides[get_session] = _get_session + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + created = await client.post( + "/api/image-generations", + data={"title": "Forkable badge", "prompt": "A minimalist badge"}, + ) + assert created.status_code == 201 + parent = created.json() + + first_edit = await client.post( + f"/api/image-generations/{parent['id']}/edits", + data={"prompt": "Make the badge blue"}, + ) + assert first_edit.status_code == 201 + + fork_edit = await client.post( + f"/api/image-generations/{parent['id']}/edits", + data={"prompt": "Make the badge green"}, + ) + assert fork_edit.status_code == 201 + + paths = lineage_paths((await client.get("/api/image-generations")).json()) + + assert sorted([[item["prompt"] for item in path] for path in paths]) == [ + ["A minimalist badge", "Make the badge blue"], + ["A minimalist badge", "Make the badge green"], + ] + + +async def test_lineage_path_can_be_deleted_leaf_to_root(tmp_path, db_sessions, services) -> None: + app = create_app(make_settings(tmp_path)) + app.state.services = services + + async def _get_session(): + async with db_sessions() as session: + yield session + + app.dependency_overrides[get_session] = _get_session + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + created = await client.post( + "/api/image-generations", + data={"title": "Forkable badge", "prompt": "A minimalist badge"}, + ) + assert created.status_code == 201 + parent = created.json() + first_edit = await client.post( + f"/api/image-generations/{parent['id']}/edits", + data={"prompt": "Make the badge blue"}, + ) + assert first_edit.status_code == 201 + fork_edit = await client.post( + f"/api/image-generations/{parent['id']}/edits", + data={"prompt": "Make the badge green"}, + ) + assert fork_edit.status_code == 201 + + path = next( + path + for path in lineage_paths((await client.get("/api/image-generations")).json()) + if path[-1]["prompt"] == "Make the badge blue" + ) + for generation in reversed(path): + response = await client.delete(f"/api/image-generations/{generation['id']}") + assert response.status_code == 204 + + paths = lineage_paths((await client.get("/api/image-generations")).json()) + + assert [[item["prompt"] for item in path] for path in paths] == [["Make the badge green"]] + + async def test_image_generation_reference_validation(tmp_path, db_sessions, services) -> None: app = create_app(make_settings(tmp_path)) app.state.services = services diff --git a/tests/unit/test_design_lineage.py b/tests/unit/test_design_lineage.py new file mode 100644 index 0000000..5d2c732 --- /dev/null +++ b/tests/unit/test_design_lineage.py @@ -0,0 +1,54 @@ +from ecommerce_agent.dashboard.design_lineage import branch_label, lineage_paths, version_label + + +def generation(id_: str, parent: str | None = None, created: int = 0) -> dict[str, object]: + return { + "id": id_, + "parent_generation_id": parent, + "created_at": created, + } + + +def test_linear_chain_has_one_lineage_path() -> None: + paths = lineage_paths( + [ + generation("v3", "v2", 3), + generation("v1", None, 1), + generation("v2", "v1", 2), + ] + ) + + assert [[item["id"] for item in path] for path in paths] == [["v1", "v2", "v3"]] + assert branch_label(paths[0]) == "v1 -> v2 -> v3" + + +def test_fork_from_middle_has_two_leaf_paths() -> None: + paths = lineage_paths( + [ + generation("v1", None, 1), + generation("v2", "v1", 2), + generation("v3a", "v2", 3), + generation("v3b", "v2", 4), + ] + ) + + assert [[item["id"] for item in path] for path in paths] == [ + ["v1", "v2", "v3b"], + ["v1", "v2", "v3a"], + ] + + +def test_version_labels_are_path_local() -> None: + paths = lineage_paths( + [ + generation("root", None, 1), + generation("shared", "root", 2), + generation("leaf-a", "shared", 3), + generation("leaf-b", "shared", 4), + ] + ) + + for path in paths: + assert version_label(path, path[0]) == "v1" + assert version_label(path, path[1]) == "v2" + assert version_label(path, path[2]) == "v3"