diff --git a/.env.example b/.env.example index 84ab152..9f7103e 100644 --- a/.env.example +++ b/.env.example @@ -20,6 +20,14 @@ ETSY_SHOP_ID= ETSY_REDIRECT_URI=http://localhost:8000/api/etsy/oauth/callback ETSY_OAUTH_STATE_TTL_SECONDS=600 +# Optional GA4 read-only source for first-party traffic validation. +GOOGLE_CLIENT_ID= +GOOGLE_CLIENT_SECRET= +# Generate once with the same Fernet command used for Etsy above. +GOOGLE_TOKEN_ENCRYPTION_KEY= +GOOGLE_ANALYTICS_REDIRECT_URI=http://localhost:8000/api/research-sources/google-analytics/oauth/callback +GOOGLE_OAUTH_STATE_TTL_SECONDS=600 + PRINTFUL_ACCESS_TOKEN= PRINTFUL_STORE_ID= PRINTFUL_ETSY_STORE_ID= @@ -48,3 +56,7 @@ OPENAI_IMAGE_WIDTH=1024 OPENAI_IMAGE_HEIGHT=1024 OPENAI_IMAGE_MAX_COST_USD=0.50 OPENAI_IMAGE_REFERENCE_COST_RESERVE_USD=0.06 + +RESEARCH_TREND_LIMIT=8 +RESEARCH_POLL_SECONDS=5 +RESEARCH_MAX_TOOL_CALLS=12 diff --git a/.githooks/pre-commit b/.githooks/pre-commit index a6f212b..faaf5ae 100755 --- a/.githooks/pre-commit +++ b/.githooks/pre-commit @@ -6,3 +6,4 @@ cd "$root" make lint make format-check +make typecheck diff --git a/README.md b/README.md index f1150d1..07298df 100644 --- a/README.md +++ b/README.md @@ -2,12 +2,12 @@ Python 3.12 MVP for a reviewable Etsy and Printful product workflow: -`trend research -> topic approval -> design -> design approval -> listing copy approval -> Printful product -> mock Etsy draft -> manual publish -> monitoring` +`trend discovery -> optional research report -> design -> design approval -> listing copy approval -> Printful product -> Etsy draft -> manual publish -> monitoring` -Etsy remains mocked. Printful supports real authenticated catalog reads and an -operator-driven Workshop for curation, mockups, and sync-product creation; all real -writes require a separate feature flag and S3-compatible asset storage. OpenAI image -generation can still be enabled independently. +Etsy supports authenticated draft-listing workflows. Printful supports real catalog +reads and an operator-driven Workshop for curation, mockups, and sync-product +creation; all real writes require separate safeguards. OpenAI image generation and +web-grounded research can be enabled independently. The architecture rationale and the decisions behind it are documented in [ARCHITECTURE_REVIEW.md](ARCHITECTURE_REVIEW.md). @@ -112,11 +112,15 @@ The MVP is local-only and has no authentication. Do not expose ports remotely. A ## Data Model -Ten tables, all created by Alembic: +Core tables created by Alembic include: | Table | Holds | |---|---| -| `trends` | One row per researched trend candidate, with embedded scores and a review status | +| `trend_discovery_runs` | Durable manual scans for emerging apparel trends | +| `trends` | Trend candidates with evidence, apparel scores, confidence, and risk warnings | +| `research_runs` | In-depth research tasks, structured output, and rendered Markdown | +| `research_source_configs` | Selected optional research sources such as GA4 | +| `etsy_stats_imports` | Deduplicated manual imports of Etsy shopper search terms | | `products` | The full concept → design → copy → listings → publish lifecycle of one product | | `design_assets` | Every generated design revision with its brief, checksum, and review outcome | | `approval_events` | Append-only human decisions with unique idempotency keys | @@ -158,6 +162,41 @@ before generation when its configured output-cost estimate exceeds the per-image limit. Dollar amounts are estimates because the generation response reports token usage rather than the final billed amount. +## Trend Research And Reports + +The Trend Inbox starts manual discovery scans. In real OpenAI mode, scans and reports +use Responses API web search in background mode, store verified source URLs, and +continue through the durable job queue while the dashboard is closed. Mock mode +provides deterministic local results. + +Each trend includes a short explanation, 1–5 apparel score, confidence, evidence, and +separate trademark, copyright, cultural, and marketplace-policy warnings. Selecting +**Research** opens an editable topic form. Manual topics use the same form. Completed +reports render consistently in the dashboard and download as Markdown. + +Research context can include Etsy API listing and transaction signals, manually +imported Etsy Stats search terms, an optional read-only GA4 property, and cached +Printful products, prices, placements, and techniques. + +To connect GA4, create a Google OAuth web client, enable the Google Analytics Admin +and Data APIs, add the callback as an authorized redirect URI, and set: + +```text +GOOGLE_CLIENT_ID=... +GOOGLE_CLIENT_SECRET=... +GOOGLE_TOKEN_ENCRYPTION_KEY=... # Fernet key +GOOGLE_ANALYTICS_REDIRECT_URI=https://api.workshop.example.com/api/research-sources/google-analytics/oauth/callback +``` + +Then connect Google Analytics and select a property from Research Reports. GA4 is a +first-party validation signal, not broad market discovery. + +Etsy shopper search terms remain in Shop Manager Stats rather than the current public +API. Research Reports explains how to open **Shop Manager → Stats → How shoppers +found you → Etsy search**, download Workshop's CSV template, and upload the copied +terms and visit counts. Google Trends is disabled because its official API remains +limited-access alpha. + ## Tests ```bash diff --git a/TODO.md b/TODO.md index d5d4e22..6b074c0 100644 --- a/TODO.md +++ b/TODO.md @@ -140,7 +140,11 @@ duplicate actions or unexpected charges. persist bytes through `AssetStore`; record model, size, usage, cost, checksum. - [x] Per-image budget, top-level kill switch, provider mode, safe transient retries, and explicit unknown-outcome handling for ambiguous timeouts. -- [x] Durable standalone Design Studio for prompt iteration and generation history. +- [x] `APP_ENV=test` rejects `OPENAI_MODE=real`, so local smoke tests cannot inherit + paid OpenAI settings from `.env` by accident. +- [x] Durable standalone Design Studio for prompt iteration, collapsed lineage, + shared design titles, v1/v2/v3 tracking, prompt previews, and configurable aspect + ratios. - [ ] Daily image-generation budget and promotion of a selected playground image into the product design queue. - [x] Promotion of Design Studio and approved workflow assets into an independent @@ -152,6 +156,30 @@ duplicate actions or unexpected charges. - [ ] Basic print-readiness checks (dimensions, transparency, DPI) and human review of every asset (already enforced by the design gate). +### Trend research and reports + +- [x] Durable manual trend-discovery runs with emerging-theme blurbs, evidence, + confidence, 1–5 apparel recommendations, and separate IP/policy warnings. +- [x] Shared editable research form for manual topics and trend-prefilled topics. +- [x] Durable background research runs using structured output, verified web sources, + bounded transient-rate-limit retries, consistent Markdown rendering, in-app display, + and download. +- [x] Live-source validation keeps only URLs recorded by OpenAI web search, rejects + results with no verified sources, and reuses completed provider responses when a + failed report is retried. +- [x] Research context from Etsy API signals, imported Etsy Stats search terms, cached + Printful products, and optional read-only GA4 data. +- [x] Etsy Stats CSV template, validation, deduplication, import history, and in-app + instructions for collecting terms from Shop Manager. +- [x] Google Analytics OAuth, token refresh, accessible-property selection, source + status, and recent-versus-prior reporting. +- [x] Google Trends integration is documented as disabled while the official API + remains limited-access alpha; it is not a research dependency. +- [ ] Review discovery quality and OpenAI cost/latency after several real shop runs; + tune the trend limit, tool-call limit, and model only from observed results. +- [ ] Add scheduled trend discovery after manual scan quality and operating cost are + understood. + ### Notifications (Slack, after the dashboard works remotely) - [ ] `BLOCKED` Single-workspace Slack app; bot token + signing secret stored securely. diff --git a/migrations/versions/0009_research_subsystem.py b/migrations/versions/0009_research_subsystem.py new file mode 100644 index 0000000..2e36cca --- /dev/null +++ b/migrations/versions/0009_research_subsystem.py @@ -0,0 +1,245 @@ +"""Add durable trend discovery and research reports. + +Revision ID: 0009 +Revises: 0008 +Create Date: 2026-06-15 +""" + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import postgresql + +revision = "0009" +down_revision = "0008" +branch_labels = None +depends_on = None + +JSON = postgresql.JSONB(astext_type=sa.Text()) + + +def _identity_columns(*, versioned: bool = False) -> list[sa.Column]: + columns = [ + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + server_default=sa.text("now()"), + nullable=False, + ), + sa.Column( + "updated_at", + sa.DateTime(timezone=True), + server_default=sa.text("now()"), + nullable=False, + ), + ] + if versioned: + columns.append(sa.Column("version", sa.Integer(), nullable=False)) + columns.append(sa.PrimaryKeyConstraint("id")) + return columns + + +def upgrade() -> None: + op.alter_column( + "oauth_credentials", + "provider", + type_=sa.String(32), + existing_type=sa.String(8), + existing_nullable=False, + ) + op.create_table( + "google_oauth_states", + sa.Column("state_digest", sa.String(64), nullable=False), + sa.Column("code_verifier_ciphertext", sa.Text(), nullable=False), + sa.Column("redirect_uri", sa.Text(), nullable=False), + sa.Column("scopes", JSON, nullable=False), + sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("consumed_at", sa.DateTime(timezone=True)), + *_identity_columns(), + sa.UniqueConstraint("state_digest"), + ) + op.create_index( + "ix_google_oauth_states_state_digest", + "google_oauth_states", + ["state_digest"], + ) + op.create_index( + "ix_google_oauth_states_expires_at", + "google_oauth_states", + ["expires_at"], + ) + + op.create_table( + "trend_discovery_runs", + sa.Column("status", sa.String(16), nullable=False), + sa.Column("provider_response_id", sa.String(255)), + sa.Column("input_snapshot", JSON, nullable=False), + sa.Column("structured_result", JSON, nullable=False), + sa.Column("error_message", sa.Text()), + sa.Column("started_at", sa.DateTime(timezone=True)), + sa.Column("finished_at", sa.DateTime(timezone=True)), + *_identity_columns(), + ) + op.create_index("ix_trend_discovery_runs_status", "trend_discovery_runs", ["status"]) + + op.add_column("trends", sa.Column("discovery_run_id", sa.Uuid())) + op.add_column( + "trends", + sa.Column("blurb", sa.Text(), server_default=sa.text("''"), nullable=False), + ) + op.add_column("trends", sa.Column("apparel_score", sa.Numeric(3, 2))) + op.add_column( + "trends", + sa.Column( + "score_breakdown", + JSON, + server_default=sa.text("'{}'::jsonb"), + nullable=False, + ), + ) + op.add_column("trends", sa.Column("confidence", sa.Numeric(3, 2))) + op.add_column( + "trends", + sa.Column( + "verified_sources", + JSON, + server_default=sa.text("'[]'::jsonb"), + nullable=False, + ), + ) + op.add_column( + "trends", + sa.Column( + "risk_warnings", + JSON, + server_default=sa.text("'[]'::jsonb"), + nullable=False, + ), + ) + op.create_foreign_key( + "fk_trends_discovery_run_id", + "trends", + "trend_discovery_runs", + ["discovery_run_id"], + ["id"], + ) + op.create_index("ix_trends_discovery_run_id", "trends", ["discovery_run_id"]) + + connection = op.get_bind() + rows = connection.execute( + sa.text( + "SELECT research_batch_id, MIN(created_at), MAX(updated_at) " + "FROM trends GROUP BY research_batch_id" + ) + ).fetchall() + for batch_id, created_at, updated_at in rows: + connection.execute( + sa.text( + "INSERT INTO trend_discovery_runs " + "(id, status, input_snapshot, structured_result, created_at, updated_at, " + "finished_at) VALUES " + "(:id, 'completed', '{\"legacy\": true}', '{}', :created_at, :updated_at, " + ":updated_at)" + ), + { + "id": batch_id, + "created_at": created_at, + "updated_at": updated_at, + }, + ) + connection.execute( + sa.text("UPDATE trends SET discovery_run_id = :id WHERE research_batch_id = :id"), + {"id": batch_id}, + ) + + op.create_table( + "research_runs", + sa.Column("trend_id", sa.Uuid()), + sa.Column("topic_title", sa.String(255), nullable=False), + sa.Column( + "additional_context", + sa.Text(), + server_default=sa.text("''"), + nullable=False, + ), + sa.Column("status", sa.String(16), nullable=False), + sa.Column("provider_response_id", sa.String(255)), + sa.Column("input_snapshot", JSON, nullable=False), + sa.Column("structured_result", JSON, nullable=False), + sa.Column("markdown_report", sa.Text()), + sa.Column("error_message", sa.Text()), + sa.Column("started_at", sa.DateTime(timezone=True)), + sa.Column("finished_at", sa.DateTime(timezone=True)), + *_identity_columns(), + sa.ForeignKeyConstraint(["trend_id"], ["trends.id"]), + ) + op.create_index("ix_research_runs_trend_id", "research_runs", ["trend_id"]) + op.create_index("ix_research_runs_status", "research_runs", ["status"]) + + op.create_table( + "research_source_configs", + sa.Column("provider", sa.String(50), nullable=False), + sa.Column("enabled", sa.Boolean(), nullable=False), + sa.Column("config", JSON, nullable=False), + *_identity_columns(versioned=True), + sa.UniqueConstraint("provider"), + ) + + op.create_table( + "etsy_stats_imports", + sa.Column("filename", sa.String(255), nullable=False), + sa.Column("checksum", sa.String(64), nullable=False), + sa.Column("period_start", sa.DateTime(timezone=True), nullable=False), + sa.Column("period_end", sa.DateTime(timezone=True), nullable=False), + sa.Column("row_count", sa.Integer(), nullable=False), + sa.Column("rows", JSON, nullable=False), + sa.Column("warnings", JSON, nullable=False), + *_identity_columns(), + sa.UniqueConstraint("checksum"), + ) + op.create_index("ix_etsy_stats_imports_checksum", "etsy_stats_imports", ["checksum"]) + + op.add_column("jobs", sa.Column("available_at", sa.DateTime(timezone=True))) + op.create_index("ix_jobs_available_at", "jobs", ["available_at"]) + op.drop_index("uq_jobs_active", table_name="jobs") + op.create_index( + "uq_jobs_active", + "jobs", + ["job_type", "subject_id"], + unique=True, + postgresql_where=sa.text("status IN ('queued', 'waiting', 'running')"), + ) + + +def downgrade() -> None: + op.drop_index("uq_jobs_active", table_name="jobs") + op.create_index( + "uq_jobs_active", + "jobs", + ["job_type", "subject_id"], + unique=True, + postgresql_where=sa.text("status IN ('queued', 'running')"), + ) + op.drop_index("ix_jobs_available_at", table_name="jobs") + op.drop_column("jobs", "available_at") + op.drop_table("etsy_stats_imports") + op.drop_table("research_source_configs") + op.drop_table("research_runs") + op.drop_index("ix_trends_discovery_run_id", table_name="trends") + op.drop_constraint("fk_trends_discovery_run_id", "trends", type_="foreignkey") + op.drop_column("trends", "risk_warnings") + op.drop_column("trends", "verified_sources") + op.drop_column("trends", "confidence") + op.drop_column("trends", "score_breakdown") + op.drop_column("trends", "apparel_score") + op.drop_column("trends", "blurb") + op.drop_column("trends", "discovery_run_id") + op.drop_table("trend_discovery_runs") + op.drop_table("google_oauth_states") + op.alter_column( + "oauth_credentials", + "provider", + type_=sa.String(8), + existing_type=sa.String(32), + existing_nullable=False, + ) diff --git a/migrations/versions/0010_image_generation_titles.py b/migrations/versions/0010_image_generation_titles.py new file mode 100644 index 0000000..a57232c --- /dev/null +++ b/migrations/versions/0010_image_generation_titles.py @@ -0,0 +1,30 @@ +"""Add image generation titles. + +Revision ID: 0010 +Revises: 0009 +Create Date: 2026-06-15 +""" + +import sqlalchemy as sa +from alembic import op + +revision = "0010" +down_revision = "0009" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "image_generations", + sa.Column( + "title", + sa.String(255), + server_default="Untitled design", + nullable=False, + ), + ) + + +def downgrade() -> None: + op.drop_column("image_generations", "title") diff --git a/src/ecommerce_agent/api/main.py b/src/ecommerce_agent/api/main.py index 4b33e96..064ced1 100644 --- a/src/ecommerce_agent/api/main.py +++ b/src/ecommerce_agent/api/main.py @@ -11,6 +11,8 @@ from ecommerce_agent.api.etsy_routes import router as etsy_router from ecommerce_agent.api.posting_routes import router as posting_router from ecommerce_agent.api.printful_routes import router as printful_router +from ecommerce_agent.api.research_routes import router as research_router +from ecommerce_agent.api.research_routes import runs_router as research_runs_router from ecommerce_agent.api.routes import router from ecommerce_agent.config import Settings, get_settings from ecommerce_agent.jobs.worker import run_worker_loop @@ -18,10 +20,12 @@ from ecommerce_agent.pipeline.monitoring import run_monitor_loop from ecommerce_agent.services import build_service_container from ecommerce_agent.services.etsy import EtsyError +from ecommerce_agent.services.google_analytics import GoogleAnalyticsError from ecommerce_agent.services.oauth import EtsyOAuthError from ecommerce_agent.services.openai_copy import OpenAICopyError from ecommerce_agent.services.openai_images import OpenAIImageError from ecommerce_agent.services.printful import PrintfulError +from ecommerce_agent.services.research import ResearchProviderError def create_app(settings: Settings | None = None) -> FastAPI: @@ -71,6 +75,8 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: app.include_router(printful_router, prefix="/api") app.include_router(etsy_router, prefix="/api") app.include_router(posting_router, prefix="/api") + app.include_router(research_runs_router, prefix="/api") + app.include_router(research_router, prefix="/api") @app.exception_handler(NotFoundError) async def not_found_handler(request: Request, error: NotFoundError) -> JSONResponse: @@ -100,6 +106,18 @@ async def openai_copy_error_handler(request: Request, error: OpenAICopyError) -> async def openai_image_error_handler(request: Request, error: OpenAIImageError) -> JSONResponse: return JSONResponse(status_code=502, content={"detail": str(error)}) + @app.exception_handler(ResearchProviderError) + async def research_error_handler( + request: Request, error: ResearchProviderError + ) -> JSONResponse: + return JSONResponse(status_code=502, content={"detail": str(error)}) + + @app.exception_handler(GoogleAnalyticsError) + async def google_analytics_error_handler( + request: Request, error: GoogleAnalyticsError + ) -> JSONResponse: + return JSONResponse(status_code=400, content={"detail": str(error)}) + @app.exception_handler(Exception) async def unexpected_error_handler(request: Request, error: Exception) -> JSONResponse: logger.error( diff --git a/src/ecommerce_agent/api/research_routes.py b/src/ecommerce_agent/api/research_routes.py new file mode 100644 index 0000000..d1e3ac6 --- /dev/null +++ b/src/ecommerce_agent/api/research_routes.py @@ -0,0 +1,451 @@ +import csv +import hashlib +import io +import re +import uuid +from datetime import UTC, datetime +from typing import Annotated, Any + +import httpx +from fastapi import APIRouter, Depends, File, Query, Request, UploadFile, status +from fastapi.responses import HTMLResponse, PlainTextResponse, RedirectResponse +from sqlalchemy import desc, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from ecommerce_agent.api.schemas import ( + GoogleAnalyticsPropertyRequest, + ResearchRunCreateRequest, + RunResponse, +) +from ecommerce_agent.api.serialization import serialize_model +from ecommerce_agent.db.models import ( + EtsyStatsImport, + OAuthCredential, + PrintfulCatalogProduct, + ResearchRun, + ResearchSourceConfig, + Trend, + TrendDiscoveryRun, +) +from ecommerce_agent.db.session import get_session +from ecommerce_agent.domain.enums import ProviderName, ResearchRunStatus +from ecommerce_agent.jobs.queue import enqueue +from ecommerce_agent.pipeline.errors import ConflictError, NotFoundError +from ecommerce_agent.services.google_analytics import ( + GoogleAnalyticsError, + GoogleTokenCipher, + complete_google_authorization, + create_google_authorization_url, + disconnect_google_analytics, + list_ga4_properties, + select_ga4_property, +) + +router = APIRouter(prefix="/research-sources", tags=["research"]) +runs_router = APIRouter(tags=["research"]) +SessionDep = Annotated[AsyncSession, Depends(get_session)] + +ETSY_STATS_HEADERS = ( + "period_start", + "period_end", + "search_term", + "visits", + "listing_id", + "listing_title", + "notes", +) + + +@runs_router.post( + "/trend-discovery-runs", + response_model=RunResponse, + status_code=status.HTTP_201_CREATED, +) +async def create_trend_discovery_run(request: Request, session: SessionDep) -> RunResponse: + limit = request.app.state.settings.research_trend_limit + run = TrendDiscoveryRun( + status=ResearchRunStatus.QUEUED, + input_snapshot={"requested_limit": limit}, + ) + session.add(run) + await session.flush() + await enqueue(session, "discover_trends", run.id) + await session.commit() + return RunResponse(id=run.id, subject_type="trend_discovery_run", status=run.status.value) + + +@runs_router.get("/trend-discovery-runs") +async def list_trend_discovery_runs(session: SessionDep, limit: int = 50) -> list[dict[str, Any]]: + runs = ( + await session.scalars( + select(TrendDiscoveryRun) + .order_by(desc(TrendDiscoveryRun.created_at)) + .limit(min(limit, 200)) + ) + ).all() + rows = [] + for run in runs: + trend_count = await session.scalar( + select(func.count()).select_from(Trend).where(Trend.discovery_run_id == run.id) + ) + rows.append({**serialize_model(run), "trend_count": trend_count or 0}) + return rows + + +@runs_router.post( + "/research-runs", + response_model=RunResponse, + status_code=status.HTTP_201_CREATED, +) +async def create_research_run(body: ResearchRunCreateRequest, session: SessionDep) -> RunResponse: + if body.trend_id is not None and await session.get(Trend, body.trend_id) is None: + raise NotFoundError("Trend not found.") + run = ResearchRun( + trend_id=body.trend_id, + topic_title=body.topic_title.strip(), + additional_context=body.additional_context.strip(), + status=ResearchRunStatus.QUEUED, + ) + session.add(run) + await session.flush() + await enqueue(session, "generate_research_report", run.id) + await session.commit() + return RunResponse(id=run.id, subject_type="research_run", status=run.status.value) + + +@runs_router.get("/research-runs") +async def list_research_runs(session: SessionDep, limit: int = 100) -> list[dict[str, Any]]: + rows = ( + await session.scalars( + select(ResearchRun).order_by(desc(ResearchRun.created_at)).limit(min(limit, 500)) + ) + ).all() + return [_research_run_view(row, include_report=False) for row in rows] + + +@runs_router.get("/research-runs/{run_id}") +async def get_research_run(run_id: uuid.UUID, session: SessionDep) -> dict[str, Any]: + run = await session.get(ResearchRun, run_id) + if run is None: + raise NotFoundError("Research run not found.") + return _research_run_view(run, include_report=True) + + +@runs_router.get("/research-runs/{run_id}/download") +async def download_research_run(run_id: uuid.UUID, session: SessionDep) -> PlainTextResponse: + run = await session.get(ResearchRun, run_id) + if run is None: + raise NotFoundError("Research run not found.") + if run.status != ResearchRunStatus.COMPLETED or not run.markdown_report: + raise ConflictError("Research report is not ready for download.") + filename = re.sub(r"[^a-z0-9]+", "-", run.topic_title.lower()).strip("-") + filename = (filename or "research-report")[:80] + return PlainTextResponse( + run.markdown_report, + media_type="text/markdown", + headers={"Content-Disposition": f'attachment; filename="{filename}.md"'}, + ) + + +@runs_router.post("/research-runs/{run_id}/retry", response_model=RunResponse) +async def retry_research_run(run_id: uuid.UUID, session: SessionDep) -> RunResponse: + run = await session.get(ResearchRun, run_id) + if run is None: + raise NotFoundError("Research run not found.") + if run.status != ResearchRunStatus.FAILED: + raise ConflictError("Only failed research runs can be retried.") + run.status = ResearchRunStatus.QUEUED + snapshot = dict(run.input_snapshot) + snapshot["_workshop_manual_retry"] = True + run.input_snapshot = snapshot + run.error_message = None + run.finished_at = None + await enqueue(session, "generate_research_report", run.id) + await session.commit() + return RunResponse(id=run.id, subject_type="research_run", status=run.status.value) + + +@router.get("/status") +async def research_source_status(request: Request, session: SessionDep) -> dict[str, Any]: + settings = request.app.state.settings + google_credential = await session.scalar( + select(OAuthCredential).where(OAuthCredential.provider == ProviderName.GOOGLE_ANALYTICS) + ) + etsy_credential = await session.scalar( + select(OAuthCredential).where(OAuthCredential.provider == ProviderName.ETSY) + ) + google_config = await session.scalar( + select(ResearchSourceConfig).where(ResearchSourceConfig.provider == "google_analytics") + ) + imports = await session.scalar(select(func.count()).select_from(EtsyStatsImport)) + printful_products = await session.scalar( + select(func.count()).select_from(PrintfulCatalogProduct) + ) + return { + "web_research": { + "enabled": settings.openai_mode == "real", + "mode": settings.openai_mode, + }, + "etsy_api": { + "enabled": settings.etsy_mode == "real", + "connected": etsy_credential is not None, + "mode": settings.etsy_mode, + }, + "etsy_stats": {"import_count": imports or 0}, + "google_analytics": { + "configured": bool( + settings.google_client_id + and settings.google_client_secret + and settings.google_token_encryption_key + ), + "connected": google_credential is not None, + "account_id": google_credential.account_id if google_credential else None, + "property": google_config.config if google_config and google_config.enabled else None, + }, + "google_trends": { + "enabled": False, + "reason": "Official Google Trends API access remains limited alpha.", + }, + "printful_catalog": {"cached_product_count": printful_products or 0}, + } + + +@router.get("/etsy-stats-imports") +async def list_etsy_stats_imports(session: SessionDep, limit: int = 50) -> list[dict[str, Any]]: + rows = ( + await session.scalars( + select(EtsyStatsImport) + .order_by(desc(EtsyStatsImport.created_at)) + .limit(min(limit, 200)) + ) + ).all() + return [ + { + "id": str(row.id), + "filename": row.filename, + "period_start": row.period_start, + "period_end": row.period_end, + "row_count": row.row_count, + "warnings": row.warnings, + "created_at": row.created_at, + } + for row in rows + ] + + +@router.get("/etsy-stats-imports/template") +async def etsy_stats_template() -> PlainTextResponse: + output = io.StringIO() + writer = csv.writer(output) + writer.writerow(ETSY_STATS_HEADERS) + writer.writerow( + [ + "2026-05-01", + "2026-05-31", + "example search term", + "12", + "", + "", + "Optional note", + ] + ) + return PlainTextResponse( + output.getvalue(), + media_type="text/csv", + headers={"Content-Disposition": 'attachment; filename="workshop-etsy-stats-template.csv"'}, + ) + + +@router.post( + "/etsy-stats-imports", + status_code=status.HTTP_201_CREATED, +) +async def import_etsy_stats( + session: SessionDep, + upload: Annotated[UploadFile, File()], +) -> dict[str, Any]: + content = await upload.read() + if len(content) > 2_000_000: + raise ConflictError("Etsy Stats CSV must be 2 MB or smaller.") + checksum = hashlib.sha256(content).hexdigest() + existing = await session.scalar( + select(EtsyStatsImport).where(EtsyStatsImport.checksum == checksum) + ) + if existing is not None: + raise ConflictError("This Etsy Stats CSV has already been imported.") + try: + text = content.decode("utf-8-sig") + except UnicodeDecodeError as error: + raise ConflictError("Etsy Stats CSV must be UTF-8 encoded.") from error + reader = csv.DictReader(io.StringIO(text)) + if reader.fieldnames is None or set(ETSY_STATS_HEADERS) - set(reader.fieldnames): + raise ConflictError("Etsy Stats CSV must contain: " + ", ".join(ETSY_STATS_HEADERS) + ".") + rows: list[dict[str, Any]] = [] + warnings: list[str] = [] + starts: list[datetime] = [] + ends: list[datetime] = [] + for line_number, raw in enumerate(reader, start=2): + if len(rows) >= 5000: + raise ConflictError("Etsy Stats CSV may contain at most 5,000 data rows.") + term = (raw.get("search_term") or "").strip() + if not term: + raise ConflictError(f"Row {line_number}: search_term is required.") + try: + period_start = _date(raw.get("period_start") or "") + period_end = _date(raw.get("period_end") or "") + visits = int((raw.get("visits") or "").strip()) + except ValueError as error: + raise ConflictError( + f"Row {line_number}: dates must be YYYY-MM-DD and visits an integer." + ) from error + if period_end < period_start: + raise ConflictError(f"Row {line_number}: period_end precedes period_start.") + if visits < 0: + raise ConflictError(f"Row {line_number}: visits cannot be negative.") + listing_id = (raw.get("listing_id") or "").strip() or None + if listing_id and not listing_id.isdigit(): + warnings.append(f"Row {line_number}: listing_id is not numeric.") + rows.append( + { + "period_start": period_start.date().isoformat(), + "period_end": period_end.date().isoformat(), + "search_term": term, + "visits": visits, + "listing_id": listing_id, + "listing_title": (raw.get("listing_title") or "").strip() or None, + "notes": (raw.get("notes") or "").strip() or None, + } + ) + starts.append(period_start) + ends.append(period_end) + if not rows: + raise ConflictError("Etsy Stats CSV contains no data rows.") + record = EtsyStatsImport( + filename=(upload.filename or "etsy-stats.csv")[:255], + checksum=checksum, + period_start=min(starts), + period_end=max(ends), + row_count=len(rows), + rows=rows, + warnings=warnings, + ) + session.add(record) + await session.commit() + return { + "id": str(record.id), + "filename": record.filename, + "row_count": record.row_count, + "warnings": record.warnings, + } + + +def _google_settings(request: Request) -> tuple[Any, GoogleTokenCipher]: + settings = request.app.state.settings + if ( + not settings.google_client_id + or settings.google_client_secret is None + or settings.google_token_encryption_key is None + ): + raise GoogleAnalyticsError("Google Analytics OAuth is not configured.") + return ( + settings, + GoogleTokenCipher(settings.google_token_encryption_key.get_secret_value()), + ) + + +@router.get("/google-analytics/oauth/start") +async def start_google_oauth(request: Request, session: SessionDep) -> RedirectResponse: + settings, cipher = _google_settings(request) + url = await create_google_authorization_url( + session, + cipher=cipher, + client_id=settings.google_client_id, + redirect_uri=settings.google_analytics_redirect_uri, + ttl_seconds=settings.google_oauth_state_ttl_seconds, + ) + await session.commit() + return RedirectResponse(url, status_code=307) + + +@router.get("/google-analytics/oauth/callback", response_class=HTMLResponse) +async def complete_google_oauth( + request: Request, + session: SessionDep, + code: str | None = Query(default=None), + state: str | None = Query(default=None), + error: str | None = Query(default=None), + error_description: str | None = Query(default=None), +) -> HTMLResponse: + if error: + raise GoogleAnalyticsError( + f"Google authorization was declined: {error_description or error}." + ) + if not code or not state: + raise GoogleAnalyticsError("Google callback is missing code or state.") + settings, cipher = _google_settings(request) + async with httpx.AsyncClient(timeout=30) as client: + credential = await complete_google_authorization( + session, + cipher=cipher, + client_id=settings.google_client_id, + client_secret=settings.google_client_secret.get_secret_value(), + state=state, + code=code, + client=client, + ) + await session.commit() + return HTMLResponse( + "

Google Analytics connected

" + f"

Authenticated as {credential.account_id}.

" + "

Return to Workshop and select a GA4 property.

" + ) + + +@router.get("/google-analytics/properties") +async def google_analytics_properties( + request: Request, session: SessionDep +) -> list[dict[str, str]]: + settings, cipher = _google_settings(request) + async with httpx.AsyncClient(timeout=30) as client: + properties = await list_ga4_properties( + session, + cipher=cipher, + client_id=settings.google_client_id, + client_secret=settings.google_client_secret.get_secret_value(), + client=client, + ) + await session.commit() + return properties + + +@router.post("/google-analytics/property") +async def set_google_analytics_property( + body: GoogleAnalyticsPropertyRequest, session: SessionDep +) -> dict[str, Any]: + config = await select_ga4_property( + session, + property_id=body.property_id, + display_name=body.display_name, + ) + await session.commit() + return config.config + + +@router.post("/google-analytics/disconnect") +async def disconnect_google(session: SessionDep) -> dict[str, bool]: + await disconnect_google_analytics(session) + await session.commit() + return {"connected": False} + + +def _research_run_view(run: ResearchRun, *, include_report: bool) -> dict[str, Any]: + payload = serialize_model(run) + if not include_report: + payload.pop("input_snapshot", None) + payload.pop("structured_result", None) + payload.pop("markdown_report", None) + return payload + + +def _date(value: str) -> datetime: + return datetime.strptime(value.strip(), "%Y-%m-%d").replace(tzinfo=UTC) diff --git a/src/ecommerce_agent/api/routes.py b/src/ecommerce_agent/api/routes.py index 10095af..948b534 100644 --- a/src/ecommerce_agent/api/routes.py +++ b/src/ecommerce_agent/api/routes.py @@ -14,6 +14,7 @@ from ecommerce_agent.api.schemas import ( DecisionRequest, ImageGenerationCreateRequest, + ImageGenerationUpdateRequest, PublishRequest, RunResponse, ) @@ -55,6 +56,7 @@ ServicesDep = Annotated[ServiceContainer, Depends(get_services)] MAX_REFERENCE_BYTES = 10 * 1024 * 1024 MAX_REFERENCE_PIXELS = 25_000_000 +MAX_IMAGE_PIXELS = 4_194_304 REFERENCE_CONTENT_TYPES = { "PNG": "image/png", "JPEG": "image/jpeg", @@ -65,6 +67,9 @@ @dataclass(frozen=True) class ParsedImageRequest: prompt: str + title: str + width: int | None + height: int | None output_format: str background_removal_friendly: bool references: tuple[ImageReference, ...] @@ -118,6 +123,9 @@ async def edit_image_generation( ) parsed = ParsedImageRequest( prompt=parsed.prompt, + title=parent.title, + width=parsed.width, + height=parsed.height, output_format=parsed.output_format, background_removal_friendly=parsed.background_removal_friendly, references=(parent_reference, *parsed.references), @@ -140,14 +148,16 @@ async def _execute_image_generation( services: ServiceContainer, ) -> dict[str, Any]: settings = request.app.state.settings + width = parsed.width or settings.openai_image_width + height = parsed.height or settings.openai_image_height is_real = settings.openai_mode == "real" try: output_estimate = ( estimate_image_output_cost( model=settings.openai_image_model, quality=settings.openai_image_quality, - width=settings.openai_image_width, - height=settings.openai_image_height, + width=width, + height=height, ) if is_real else Decimal("0") @@ -163,10 +173,11 @@ async def _execute_image_generation( generation = ImageGeneration( parent_generation_id=parent_generation_id, + title=parsed.title, prompt=parsed.prompt, status=ImageGenerationStatus.PENDING, - width=settings.openai_image_width, - height=settings.openai_image_height, + width=width, + height=height, provider="openai" if is_real else "mock", model=settings.openai_image_model if is_real else "mock-image-v1", quality=settings.openai_image_quality, @@ -222,11 +233,19 @@ async def _execute_image_generation( async def _parse_image_request(request: Request, *, max_references: int) -> ParsedImageRequest: content_type = request.headers.get("content-type", "") + prompt_value: Any + title_value: Any + width_value: Any + height_value: Any + output_format_value: Any if content_type.startswith("multipart/form-data") or content_type.startswith( "application/x-www-form-urlencoded" ): form = await request.form() prompt_value = form.get("prompt") + title_value = form.get("title", "Untitled design") + width_value = form.get("width") + height_value = form.get("height") output_format_value = form.get("output_format", "png") background_value: Any = form.get("background_removal_friendly", "true") uploads = [ @@ -248,6 +267,9 @@ async def _parse_image_request(request: Request, *, max_references: int) -> Pars status_code=422, detail="Invalid image generation request." ) from error prompt_value = body.prompt + title_value = body.title + width_value = body.width + height_value = body.height output_format_value = body.output_format background_value = body.background_removal_friendly references = () @@ -258,6 +280,15 @@ async def _parse_image_request(request: Request, *, max_references: int) -> Pars status_code=422, detail="Prompt must contain between 1 and 32000 characters.", ) + title = str(title_value or "Untitled design").strip() or "Untitled design" + if len(title) > 255: + raise HTTPException(status_code=422, detail="Title must be 255 characters or fewer.") + width = _parse_dimension(width_value, "width") + height = _parse_dimension(height_value, "height") + if (width is None) != (height is None): + raise HTTPException(status_code=422, detail="Width and height must be supplied together.") + if width is not None and height is not None and width * height > MAX_IMAGE_PIXELS: + raise HTTPException(status_code=422, detail="Image dimensions are too large.") output_format = str(output_format_value).lower() if output_format not in SUPPORTED_OUTPUT_FORMATS: raise HTTPException( @@ -267,6 +298,9 @@ async def _parse_image_request(request: Request, *, max_references: int) -> Pars background_removal_friendly = _parse_bool(background_value) return ParsedImageRequest( prompt=prompt, + title=title, + width=width, + height=height, output_format=output_format, background_removal_friendly=background_removal_friendly, references=references, @@ -328,6 +362,18 @@ def _parse_bool(value: Any) -> bool: raise HTTPException(status_code=422, detail="Invalid background-removal option.") +def _parse_dimension(value: Any, label: str) -> int | None: + if value in (None, ""): + return None + try: + dimension = int(value) + except (TypeError, ValueError) as error: + raise HTTPException(status_code=422, detail=f"Invalid image {label}.") from error + if dimension <= 0 or dimension > 4096: + raise HTTPException(status_code=422, detail=f"Invalid image {label}.") + return dimension + + def _extension(output_format: str) -> str: return "jpg" if output_format == "jpeg" else output_format @@ -348,6 +394,21 @@ async def list_image_generations(session: SessionDep, limit: int = 100) -> list[ return [serialize_model(row) for row in rows] +@router.patch("/image-generations/{generation_id}") +async def update_image_generation( + generation_id: uuid.UUID, + body: ImageGenerationUpdateRequest, + session: SessionDep, +) -> dict[str, Any]: + generation = await session.get(ImageGeneration, generation_id) + if generation is None: + raise NotFoundError("Image generation not found.") + generation.title = body.title.strip() + await session.commit() + await session.refresh(generation) + return serialize_model(generation) + + @router.delete( "/image-generations/{generation_id}", status_code=status.HTTP_204_NO_CONTENT, @@ -475,7 +536,13 @@ async def monitor( @router.get("/trends") async def list_trends(session: SessionDep, limit: int = 100) -> list[dict[str, Any]]: - rows = await Repository(session).list(Trend, limit=min(limit, 500)) + rows = ( + await session.scalars( + select(Trend) + .order_by(Trend.created_at.desc(), Trend.position.asc()) + .limit(min(limit, 500)) + ) + ).all() return [serialize_model(row) for row in rows] diff --git a/src/ecommerce_agent/api/schemas.py b/src/ecommerce_agent/api/schemas.py index 8d99df7..f725456 100644 --- a/src/ecommerce_agent/api/schemas.py +++ b/src/ecommerce_agent/api/schemas.py @@ -28,12 +28,30 @@ class RunResponse(BaseModel): status: str +class ResearchRunCreateRequest(BaseModel): + topic_title: str = Field(min_length=1, max_length=255) + additional_context: str = Field(default="", max_length=12000) + trend_id: uuid.UUID | None = None + + +class GoogleAnalyticsPropertyRequest(BaseModel): + property_id: str = Field(min_length=1, max_length=100) + display_name: str = Field(min_length=1, max_length=255) + + class ImageGenerationCreateRequest(BaseModel): prompt: str = Field(min_length=1, max_length=32000) + title: str = Field(default="Untitled design", max_length=255) + width: int | None = Field(default=None, gt=0, le=4096) + height: int | None = Field(default=None, gt=0, le=4096) output_format: str = "png" background_removal_friendly: bool = True +class ImageGenerationUpdateRequest(BaseModel): + title: str = Field(min_length=1, max_length=255) + + class PrintfulCatalogRefreshRequest(BaseModel): category_ids: list[int] | None = None diff --git a/src/ecommerce_agent/config.py b/src/ecommerce_agent/config.py index 48b1e81..f58b4f4 100644 --- a/src/ecommerce_agent/config.py +++ b/src/ecommerce_agent/config.py @@ -35,6 +35,13 @@ class Settings(BaseSettings): etsy_shop_id: str | None = None etsy_redirect_uri: str = "http://localhost:8000/api/etsy/oauth/callback" etsy_oauth_state_ttl_seconds: int = Field(default=600, ge=60, le=3600) + google_client_id: str | None = None + google_client_secret: SecretStr | None = None + google_token_encryption_key: SecretStr | None = None + google_analytics_redirect_uri: str = ( + "http://localhost:8000/api/research-sources/google-analytics/oauth/callback" + ) + google_oauth_state_ttl_seconds: int = Field(default=600, ge=60, le=3600) printful_etsy_store_id: str | None = None printful_access_token: str | None = None printful_store_id: str | None = None @@ -56,9 +63,14 @@ class Settings(BaseSettings): openai_image_height: int = Field(default=1024, gt=0) openai_image_max_cost_usd: float = Field(default=0.50, gt=0) openai_image_reference_cost_reserve_usd: float = Field(default=0.06, ge=0) + research_trend_limit: int = Field(default=8, ge=1, le=25) + research_poll_seconds: int = Field(default=5, ge=1, le=60) + research_max_tool_calls: int = Field(default=12, ge=1, le=50) @model_validator(mode="after") def validate_provider_modes(self) -> "Settings": + if self.app_env == "test" and self.openai_mode == "real": + raise ValueError("OPENAI_MODE=real is not allowed when APP_ENV=test.") if self.printful_mode == "mock" and self.app_env != "test": raise ValueError("PRINTFUL_MODE=mock is only allowed when APP_ENV=test.") real_modes = [ @@ -122,6 +134,21 @@ def validate_provider_modes(self) -> "Settings": ) if self.openai_mode == "real" and self.openai_api_key is None: raise ValueError("OPENAI_API_KEY is required when OPENAI_MODE=real.") + google_values = ( + self.google_client_id, + self.google_client_secret, + self.google_token_encryption_key, + ) + if any(value is not None for value in google_values) and not all( + value is not None for value in google_values + ): + raise ValueError( + "GOOGLE_CLIENT_ID, GOOGLE_CLIENT_SECRET, and GOOGLE_TOKEN_ENCRYPTION_KEY " + "must be configured together." + ) + parsed_google_redirect = urlparse(self.google_analytics_redirect_uri) + if not parsed_google_redirect.scheme or not parsed_google_redirect.netloc: + raise ValueError("GOOGLE_ANALYTICS_REDIRECT_URI must be an absolute URL.") return self diff --git a/src/ecommerce_agent/dashboard/app.py b/src/ecommerce_agent/dashboard/app.py index 7a09c64..c978186 100644 --- a/src/ecommerce_agent/dashboard/app.py +++ b/src/ecommerce_agent/dashboard/app.py @@ -18,10 +18,10 @@ col1, col2 = st.columns(2) with col1: - if st.button("Start mock trend research", type="primary"): + if st.button("Start trend research", type="primary"): try: - run = post("/workflows/research") - st.success(f"Research run {run['id']} is waiting for approval.") + run = post("/trend-discovery-runs") + st.success(f"Trend discovery {run['id']} was queued.") except ApiError as error: st.error(str(error)) with col2: diff --git a/src/ecommerce_agent/dashboard/client.py b/src/ecommerce_agent/dashboard/client.py index 513b746..a5798de 100644 --- a/src/ecommerce_agent/dashboard/client.py +++ b/src/ecommerce_agent/dashboard/client.py @@ -44,6 +44,17 @@ def get(path: str) -> Any: return request("GET", path) +def get_text(path: str) -> str: + try: + response = httpx.get(f"{API_BASE_URL}/api{path}", timeout=20) + response.raise_for_status() + except httpx.HTTPStatusError as error: + raise ApiError(_error_detail(error.response)) from error + except httpx.HTTPError as error: + raise ApiError(f"API is unavailable: {error}") from error + return response.text + + def post(path: str, payload: dict[str, Any] | None = None, *, timeout: float = 20) -> Any: return request("POST", path, json=payload or {}, timeout=timeout) 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/1_Trend_Inbox.py b/src/ecommerce_agent/dashboard/pages/1_Trend_Inbox.py index 0586acc..f342b62 100644 --- a/src/ecommerce_agent/dashboard/pages/1_Trend_Inbox.py +++ b/src/ecommerce_agent/dashboard/pages/1_Trend_Inbox.py @@ -1,14 +1,98 @@ import streamlit as st from ecommerce_agent.dashboard.client import get, post -from ecommerce_agent.dashboard.ui import page_header, show_error, table +from ecommerce_agent.dashboard.ui import page_header, show_error -page_header("Trend Inbox", "Discovered themes and their current workflow states.") +page_header( + "Trend Inbox", + "Discover emerging apparel themes, inspect the evidence, and start deeper research.", +) try: - if st.button("Run mock trend research", type="primary"): - post("/workflows/research") - st.rerun() - table(get("/trends"), empty="Start a research run to populate the inbox.") + sources = get("/research-sources/status") + with st.expander("Research source status"): + columns = st.columns(5) + columns[0].metric( + "Web research", + "Live" if sources["web_research"]["enabled"] else "Mock", + ) + columns[1].metric( + "Etsy API", + ( + "Connected" + if sources["etsy_api"]["connected"] + else ("Configured" if sources["etsy_api"]["enabled"] else "Mock") + ), + ) + columns[2].metric("Etsy Stats imports", sources["etsy_stats"]["import_count"]) + columns[3].metric( + "GA4", + "Connected" if sources["google_analytics"]["connected"] else "Optional", + ) + columns[4].metric( + "Printful products", + sources["printful_catalog"]["cached_product_count"], + ) + st.caption( + "Google Trends is not enabled because its official API remains limited-access alpha." + ) + + action, refresh = st.columns([1, 1]) + with action: + if st.button("Run trend research", type="primary"): + run = post("/trend-discovery-runs") + st.success(f"Trend discovery {run['id']} was queued.") + st.rerun() + with refresh: + st.button("Refresh") + + runs = get("/trend-discovery-runs") + if runs: + latest = runs[0] + status = latest["status"].replace("_", " ").title() + st.caption( + f"Latest discovery: {status} · {latest['trend_count']} trend(s) · " + f"{latest['created_at']}" + ) + if latest.get("error_message"): + st.error(latest["error_message"]) + + trends = get("/trends") + discovered = [trend for trend in trends if trend["status"] == "discovered"] + if not discovered: + st.info("Run trend research to populate the inbox.") + for trend in discovered: + with st.container(border=True): + heading, score, confidence = st.columns([5, 1, 1]) + heading.subheader(trend["title"]) + heading.caption(trend["niche"]) + score.metric( + "Apparel", + f"{float(trend['apparel_score']):.1f}/5" + if trend.get("apparel_score") is not None + else "—", + ) + confidence.metric( + "Confidence", + f"{float(trend['confidence']):.0%}" if trend.get("confidence") is not None else "—", + ) + st.write(trend.get("blurb") or "") + if trend.get("evidence"): + st.markdown("**Evidence**") + for evidence in trend["evidence"]: + st.markdown(f"- {evidence}") + if trend.get("risk_warnings"): + st.warning("\n\n".join(trend["risk_warnings"])) + sources_for_trend = trend.get("verified_sources") or [] + if sources_for_trend: + st.markdown( + "**Sources:** " + + " · ".join( + f"[{source['title']}]({source['url']})" for source in sources_for_trend + ) + ) + if st.button("Research", key=f"research-{trend['id']}"): + st.query_params["trend_id"] = trend["id"] + st.switch_page("pages/2_Research_Reports.py") except Exception as error: show_error(error) diff --git a/src/ecommerce_agent/dashboard/pages/2_Research_Reports.py b/src/ecommerce_agent/dashboard/pages/2_Research_Reports.py index debdb2b..4473881 100644 --- a/src/ecommerce_agent/dashboard/pages/2_Research_Reports.py +++ b/src/ecommerce_agent/dashboard/pages/2_Research_Reports.py @@ -1,8 +1,189 @@ -from ecommerce_agent.dashboard.client import get -from ecommerce_agent.dashboard.ui import page_header, show_error, table +import streamlit as st + +from ecommerce_agent.dashboard.client import ( + API_BASE_URL, + get, + get_text, + post, + post_multipart, +) +from ecommerce_agent.dashboard.ui import page_header, show_error + +page_header( + "Research Reports", + "Run in-depth topic research and review consistent, downloadable Markdown reports.", +) -page_header("Research Reports", "Demand, competition, pricing, and scalability evidence.") try: - table(get("/research-reports")) + trends = get("/trends") + trend_id = st.query_params.get("trend_id") + selected_trend = next( + (trend for trend in trends if trend["id"] == trend_id), + None, + ) + + with st.expander( + "Start research" + (f": {selected_trend['title']}" if selected_trend else ""), + expanded=bool(selected_trend), + ): + with st.form("research-topic-form"): + topic_title = st.text_input( + "Topic title", + value=selected_trend["title"] if selected_trend else "", + ) + default_context = "" + if selected_trend: + default_context = ( + f"Trend summary: {selected_trend.get('blurb', '')}\n\n" + f"Niche: {selected_trend.get('niche', '')}\n\n" + "Please validate the current momentum and focus on wearable, differentiated " + "ideas suitable for hats, shirts, and hoodies." + ) + additional_context = st.text_area( + "Additional context", + value=default_context, + height=180, + help="Optional guidance about audience, style, products, exclusions, or timing.", + ) + if st.form_submit_button("Submit", type="primary"): + run = post( + "/research-runs", + { + "topic_title": topic_title, + "additional_context": additional_context, + "trend_id": selected_trend["id"] if selected_trend else None, + }, + ) + st.query_params.clear() + st.query_params["run_id"] = run["id"] + st.success("Research run queued.") + st.rerun() + + with st.expander("Etsy Stats search-term import"): + st.markdown( + """ +Etsy keeps shopper search terms inside Shop Manager rather than its current public API. + +1. Open **Etsy.com → Shop Manager → Stats**. +2. Choose the date range you want to analyze. +3. Under **How shoppers found you**, select **Etsy search**. +4. Copy each visible search term and its visit count into the Workshop CSV template. +5. Add a listing ID or title when you know which listing the term belongs to, then upload it here. + +Use the same date range on every row from one export. Workshop keeps import history and rejects +an identical file if it has already been uploaded. +""" + ) + template = get_text("/research-sources/etsy-stats-imports/template") + st.download_button( + "Download Etsy Stats CSV template", + data=template, + file_name="workshop-etsy-stats-template.csv", + mime="text/csv", + ) + upload = st.file_uploader("Upload completed CSV", type=["csv"]) + if upload is not None and st.button("Import Etsy Stats"): + result = post_multipart( + "/research-sources/etsy-stats-imports", + data={}, + files=[ + ( + "upload", + ( + upload.name, + upload.getvalue(), + "text/csv", + ), + ) + ], + ) + st.success(f"Imported {result['row_count']} Etsy search-term rows.") + if result["warnings"]: + st.warning("\n".join(result["warnings"])) + imports = get("/research-sources/etsy-stats-imports") + if imports: + st.dataframe(imports, use_container_width=True, hide_index=True) + + with st.expander("Google Analytics 4"): + source_status = get("/research-sources/status")["google_analytics"] + if not source_status["configured"]: + st.info( + "Set GOOGLE_CLIENT_ID, GOOGLE_CLIENT_SECRET, and " + "GOOGLE_TOKEN_ENCRYPTION_KEY to enable GA4." + ) + elif not source_status["connected"]: + st.link_button( + "Connect Google Analytics", + f"{API_BASE_URL}/api/research-sources/google-analytics/oauth/start", + ) + else: + st.success(f"Connected as {source_status['account_id']}.") + properties = get("/research-sources/google-analytics/properties") + if properties: + labels = { + f"{item['display_name']} · {item['property_id']}": item for item in properties + } + selected_label = st.selectbox("GA4 property", list(labels)) + selected_property = labels[selected_label] + if st.button("Use this GA4 property"): + post( + "/research-sources/google-analytics/property", + selected_property, + ) + st.success("GA4 property saved.") + st.rerun() + else: + st.warning("No accessible GA4 properties were returned.") + if st.button("Disconnect Google Analytics"): + post("/research-sources/google-analytics/disconnect") + st.rerun() + + st.divider() + heading, refresh = st.columns([4, 1]) + heading.subheader("Research history") + refresh.button("Refresh") + runs = get("/research-runs") + if not runs: + st.info("No research runs yet.") + else: + run_labels = { + f"{run['topic_title']} · {run['status'].replace('_', ' ')} · {run['created_at']}": run + for run in runs + } + requested_run_id = st.query_params.get("run_id") + default_index = next( + ( + index + for index, run in enumerate(run_labels.values()) + if run["id"] == requested_run_id + ), + 0, + ) + selected_label = st.selectbox( + "Research run", + list(run_labels), + index=default_index, + ) + selected_run = get(f"/research-runs/{run_labels[selected_label]['id']}") + status_value = selected_run["status"].replace("_", " ").title() + st.caption( + f"Status: {status_value} · Created {selected_run['created_at']} · " + f"Run `{selected_run['id']}`" + ) + if selected_run.get("error_message"): + st.error(selected_run["error_message"]) + if st.button("Retry research run"): + post(f"/research-runs/{selected_run['id']}/retry") + st.rerun() + elif selected_run["status"] in {"queued", "running"}: + st.info("Research is running in the background. Use Refresh to check progress.") + elif selected_run.get("markdown_report"): + st.download_button( + "Download Markdown", + data=selected_run["markdown_report"], + file_name=f"{selected_run['topic_title']}.md", + mime="text/markdown", + ) + st.markdown(selected_run["markdown_report"]) except Exception as error: show_error(error) diff --git a/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py b/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py index 03e81f4..9bd4526 100644 --- a/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py +++ b/src/ecommerce_agent/dashboard/pages/4_Design_Studio.py @@ -1,9 +1,13 @@ -from typing import Any +from math import gcd +from typing import Any, Literal import streamlit as st -from ecommerce_agent.dashboard.client import delete, get, get_asset, post_multipart +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 OUTPUT_FORMATS = { "PNG": "png", @@ -15,29 +19,41 @@ "jpeg": ("jpg", "image/jpeg"), "webp": ("webp", "image/webp"), } +ASPECT_OPTIONS = ("Square 1:1", "Portrait 4:5", "Landscape 2:1", "Custom") +EDIT_ASPECT_OPTIONS = ("Same as source", *ASPECT_OPTIONS) page_header( "Design Studio", - "Generate and iteratively edit standalone image candidates. " - "These are not connected to the design queue yet.", + "Create standalone image designs, edit versions, and track each design line.", ) -if "design_prompt" not in st.session_state: - st.session_state.design_prompt = "" -if "design_output_format" not in st.session_state: - st.session_state.design_output_format = "PNG" -if "design_background_friendly" not in st.session_state: - st.session_state.design_background_friendly = True -if "editing_generation_id" not in st.session_state: - st.session_state.editing_generation_id = None -if "deleting_generation_id" not in st.session_state: - st.session_state.deleting_generation_id = None +DEFAULT_STATE = { + "design_title": "Untitled design", + "design_prompt": "", + "design_output_format": "PNG", + "design_background_friendly": True, + "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: + st.session_state[key] = value def reuse_generation(generation: dict[str, Any]) -> None: + st.session_state.design_title = generation.get("title") or "Untitled design" st.session_state.design_prompt = generation["prompt"] st.session_state.design_output_format = _format_label(generation["output_format"]) st.session_state.design_background_friendly = generation["background_removal_friendly"] + st.session_state.design_aspect = _aspect_label(generation["width"], generation["height"]) + st.session_state.design_custom_width = generation["width"] + st.session_state.design_custom_height = generation["height"] def start_edit(generation_id: str) -> None: @@ -63,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 [ ( @@ -87,230 +119,529 @@ def show_upload_previews(uploaded: list[Any]) -> None: st.image(item.getvalue(), caption=f"Reference {index + 1}", width=140) -def show_generated_image(asset_uri: str) -> bytes: +def show_generated_image( + asset_uri: str, *, width: Literal["stretch", "content"] | int = "stretch" +) -> bytes: content = get_asset(asset_uri) - st.image(content, width="stretch") + st.image(content, width=width) return content +def line_title(path: list[dict[str, Any]]) -> str: + leaf = path[-1] + return str(leaf.get("title") or "Untitled design") + + +def save_line_title(path: list[dict[str, Any]], title: str) -> None: + clean_title = title.strip() or "Untitled design" + patch(f"/image-generations/{path[-1]['id']}", {"title": clean_title}) + + +def aspect_dimensions( + choice: str, + custom_width: int, + custom_height: int, + *, + source: dict[str, Any] | None = None, +) -> tuple[int, int]: + if choice == "Same as source" and source is not None: + return int(source["width"]), int(source["height"]) + if choice == "Square 1:1": + return 1024, 1024 + if choice == "Portrait 4:5": + return 1024, 1280 + if choice == "Landscape 2:1": + return 1536, 768 + ratio = max(custom_width, 1) / max(custom_height, 1) + if ratio >= 1: + width = 1536 + height = int(round(width / ratio)) + else: + height = 1536 + width = int(round(height * ratio)) + return _round_dimension(width), _round_dimension(height) + + +def _round_dimension(value: int) -> int: + return max(64, min(4096, round(value / 64) * 64)) + + +def _aspect_label(width: int, height: int) -> str: + ratio = round(width / height, 2) + if ratio == 1: + return "Square 1:1" + if ratio == 0.8: + return "Portrait 4:5" + if ratio == 2: + return "Landscape 2:1" + return "Custom" + + +def ratio_parts(width: int, height: int) -> tuple[int, int]: + divisor = gcd(width, height) + return min(width // divisor, 20), min(height // divisor, 20) + + +def prompt_preview( + *, + prompt: str, + width: int, + height: int, + output_format: str, + background_removal_friendly: bool, + reference_count: int, +) -> str: + references = tuple( + ImageReference(f"reference-{index}.png", "image/png", b"") + for index in range(1, reference_count + 1) + ) + effective_prompt = build_effective_prompt( + ImageGenerationRequest( + prompt=prompt, + width=width, + height=height, + references=references, + output_format=output_format, + background_removal_friendly=background_removal_friendly, + product_concept_id="preview", + ) + ) + return ( + f"size: {width}x{height}\n" + f"output_format: {output_format}\n" + f"reference_images: {reference_count}\n\n" + f"prompt:\n{effective_prompt}" + ) + + def _format_label(output_format: str) -> str: return {"png": "PNG", "jpeg": "JPEG", "webp": "WebP"}[output_format] -try: - prompt = st.text_area( - "Prompt", - key="design_prompt", - height=180, - placeholder="Describe the image you want to generate...", - ) - initial_references = st.file_uploader( - "Reference images", - type=["png", "jpg", "jpeg", "webp"], - accept_multiple_files=True, - help="Upload up to four images. Upload order determines reference numbering.", +def show_details(generation: dict[str, Any]) -> None: + 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']}" ) - show_upload_previews(initial_references) + 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"]) - control_left, control_right = st.columns(2) - with control_left: - output_label = st.selectbox( - "Output format", - options=list(OUTPUT_FORMATS), - key="design_output_format", - ) - with control_right: - background_friendly = st.checkbox( - "Background-removal friendly", - key="design_background_friendly", - help="Requests an isolated subject on a uniform, high-contrast background.", + +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]: + 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"), + ): + 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.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"]) - too_many_initial = len(initial_references) > 4 - if too_many_initial: - st.error("A maximum of four reference images is allowed.") - if st.button( - "Generate image", - type="primary", - disabled=not prompt.strip() or too_many_initial, - ): - with st.spinner("Generating with OpenAI..."): - result = post_multipart( - "/image-generations", - data={ - "prompt": prompt.strip(), - "output_format": OUTPUT_FORMATS[output_label], - "background_removal_friendly": str(background_friendly).lower(), - }, - files=upload_files(initial_references), - timeout=150, - ) - if result["status"] == "succeeded": - st.success("Image generated and saved.") - elif result["status"] == "unknown": - st.warning(result["error_message"]) + 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: + if st.button( + "Confirm delete", + 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.error(result["error_message"]) - st.rerun() + 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) - generations = get("/image-generations") selected = next( - (item for item in generations if item["id"] == st.session_state.editing_generation_id), - None, + (item for item in path if item["id"] == st.session_state.focused_generation_id), + path[-1], ) - if selected is not None: - st.divider() - st.subheader("Edit design") - edit_left, edit_right = st.columns([1, 1]) - with edit_left: + 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"]) - st.caption(f"Editing render {selected['id'][:8]}") - 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...", + 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']}", ) - 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.", + with render: + render_edit = st.button( + f"Generate from {label}", + type="primary", + disabled=not edit_instruction.strip() or too_many_edit, ) - show_upload_previews(edit_references) - edit_format_label = st.selectbox( - "Edited output format", + with cancel: + 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: + generations = get("/image-generations") + + with st.expander("New design", expanded=not generations): + title = st.text_input("Title", key="design_title") + prompt = st.text_area( + "Prompt", + key="design_prompt", + height=180, + placeholder="Describe the image you want to generate...", + ) + initial_references = st.file_uploader( + "Reference images", + type=["png", "jpg", "jpeg", "webp"], + accept_multiple_files=True, + help="Upload up to four images. Upload order determines reference numbering.", + ) + show_upload_previews(initial_references) + + aspect = st.radio("Aspect ratio", ASPECT_OPTIONS, key="design_aspect", horizontal=True) + if aspect == "Custom": + custom_left, custom_right = st.columns(2) + with custom_left: + custom_width = st.number_input( + "Ratio width", min_value=1, max_value=20, key="design_custom_width" + ) + with custom_right: + custom_height = st.number_input( + "Ratio height", min_value=1, max_value=20, key="design_custom_height" + ) + else: + custom_width = st.session_state.design_custom_width + custom_height = st.session_state.design_custom_height + width, height = aspect_dimensions(aspect, custom_width, custom_height) + st.caption(f"Output size: {width}x{height}") + + control_left, control_right = st.columns(2) + with control_left: + output_label = st.selectbox( + "Output format", options=list(OUTPUT_FORMATS), - index=list(OUTPUT_FORMATS).index(_format_label(selected["output_format"])), - key=f"edit-format-{selected['id']}", + key="design_output_format", ) - edit_background = st.checkbox( + with control_right: + background_friendly = st.checkbox( "Background-removal friendly", - value=selected["background_removal_friendly"], - key=f"edit-background-{selected['id']}", + key="design_background_friendly", + help="Requests an isolated subject on a uniform, high-contrast background.", + ) + + too_many_initial = len(initial_references) > 4 + if too_many_initial: + st.error("A maximum of four reference images is allowed.") + preview, generate = st.columns([1, 1]) + with preview: + show_preview = st.button( + "Preview full prompt", + disabled=not prompt.strip(), + key="preview-create-prompt", ) - too_many_edit = len(edit_references) > 3 - if too_many_edit: - st.error("A maximum of three additional reference images is allowed.") - render, cancel = st.columns(2) - with render: - render_edit = st.button( - "Render edit", - type="primary", - disabled=not edit_instruction.strip() or too_many_edit, + with generate: + render_image = st.button( + "Generate v1", + type="primary", + disabled=not prompt.strip() or too_many_initial, + ) + if show_preview: + st.code( + prompt_preview( + prompt=prompt.strip(), + width=width, + height=height, + output_format=OUTPUT_FORMATS[output_label], + background_removal_friendly=background_friendly, + reference_count=len(initial_references), ) - with cancel: - st.button("Cancel", on_click=cancel_edit) - if render_edit: - with st.spinner("Rendering edit with OpenAI..."): - result = post_multipart( - f"/image-generations/{selected['id']}/edits", - data={ - "prompt": edit_instruction.strip(), - "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() + ) + if render_image: + with st.spinner("Generating with OpenAI..."): + result = post_multipart( + "/image-generations", + data={ + "title": title.strip() or "Untitled design", + "prompt": prompt.strip(), + "width": str(width), + "height": str(height), + "output_format": OUTPUT_FORMATS[output_label], + "background_removal_friendly": str(background_friendly).lower(), + }, + files=upload_files(initial_references), + timeout=150, + ) + if result["status"] == "succeeded": + st.success("Image generated and saved.") + elif result["status"] == "unknown": + st.warning(result["error_message"]) + else: + 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: + selected_path = next((path for path in paths if selected in path), None) + edit_dialog(selected, selected_path) st.divider() - st.subheader("Generation history") + st.subheader("Design Library") if not generations: - st.info("No playground images have been generated.") - for generation in generations: + st.info("No designs have been generated.") + + 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): - left, right = st.columns([1, 1]) - with left: - 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"design-{generation['id'][:8]}.{extension}", - mime=content_type, - key=f"download-{generation['id']}", - on_click="ignore", - width="stretch", - ) - else: - st.caption(f"Generation status: {generation['status']}") - with right: - if generation.get("parent_generation_id"): - st.caption( - f"Edit of {generation['parent_generation_id'][:8]} " - f"· version {generation['id'][:8]}" - ) - else: - st.caption(f"Original render {generation['id'][:8]}") - st.write(generation["prompt"]) - st.caption( - f"{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"): - with st.container(key=f"generation-metadata-{generation['id']}"): - st.json(generation["usage"], expanded=False) - if generation.get("error_message"): - st.error(generation["error_message"]) - actions = st.columns(3) - 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']}", - disabled=generation["status"] != "succeeded" - or not generation.get("asset_uri"), - on_click=start_edit, - args=(generation["id"],), - ) - with actions[2]: - st.button( - "Delete", - key=f"delete-{generation['id']}", - on_click=start_delete, - args=(generation["id"],), - ) - if st.session_state.deleting_generation_id == generation["id"]: - st.warning( - "Delete this design? Its saved image will also be removed. " - "Edited descendants will be preserved." - ) - confirm, cancel = st.columns(2) - with confirm: - st.button( - "Confirm delete", - key=f"confirm-delete-{generation['id']}", - type="primary", - on_click=confirm_delete, - args=(generation["id"],), - ) - with cancel: - st.button( - "Cancel", - key=f"cancel-delete-{generation['id']}", - on_click=cancel_delete, - ) + 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() + 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/src/ecommerce_agent/db/models.py b/src/ecommerce_agent/db/models.py index 20646f1..2bfec6e 100644 --- a/src/ecommerce_agent/db/models.py +++ b/src/ecommerce_agent/db/models.py @@ -40,6 +40,7 @@ PrintfulSyncStatus, ProductStatus, ProviderName, + ResearchRunStatus, TrendStatus, ) @@ -58,6 +59,9 @@ def enum_type(enum_cls: type[StrEnum], name: str) -> Enum: class Trend(UUIDPrimaryKeyMixin, TimestampMixin, VersionMixin, Base): __tablename__ = "trends" + discovery_run_id: Mapped[uuid.UUID | None] = mapped_column( + ForeignKey("trend_discovery_runs.id"), index=True + ) research_batch_id: Mapped[uuid.UUID] = mapped_column(index=True) position: Mapped[int] = mapped_column(Integer, default=0) title: Mapped[str] = mapped_column(String(255), nullable=False) @@ -65,10 +69,78 @@ class Trend(UUIDPrimaryKeyMixin, TimestampMixin, VersionMixin, Base): source: Mapped[str] = mapped_column(String(100), nullable=False) evidence: Mapped[list[str]] = mapped_column(JsonType, default=list, nullable=False) research: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict, nullable=False) + blurb: Mapped[str] = mapped_column(Text, default="", nullable=False) + apparel_score: Mapped[Decimal | None] = mapped_column(Numeric(3, 2)) + score_breakdown: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + confidence: Mapped[Decimal | None] = mapped_column(Numeric(3, 2)) + verified_sources: Mapped[list[dict[str, Any]]] = mapped_column(JsonType, default=list) + risk_warnings: Mapped[list[str]] = mapped_column(JsonType, default=list) status: Mapped[TrendStatus] = mapped_column( enum_type(TrendStatus, "trend_status"), default=TrendStatus.PENDING_REVIEW, index=True ) products: Mapped[list["Product"]] = relationship(back_populates="trend") + discovery_run: Mapped["TrendDiscoveryRun | None"] = relationship(back_populates="trends") + research_runs: Mapped[list["ResearchRun"]] = relationship(back_populates="trend") + + +class TrendDiscoveryRun(UUIDPrimaryKeyMixin, TimestampMixin, Base): + __tablename__ = "trend_discovery_runs" + + status: Mapped[ResearchRunStatus] = mapped_column( + enum_type(ResearchRunStatus, "trend_discovery_run_status"), + default=ResearchRunStatus.QUEUED, + index=True, + ) + provider_response_id: Mapped[str | None] = mapped_column(String(255)) + input_snapshot: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + structured_result: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + error_message: Mapped[str | None] = mapped_column(Text) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + trends: Mapped[list[Trend]] = relationship(back_populates="discovery_run") + + +class ResearchRun(UUIDPrimaryKeyMixin, TimestampMixin, Base): + __tablename__ = "research_runs" + + trend_id: Mapped[uuid.UUID | None] = mapped_column(ForeignKey("trends.id"), index=True) + topic_title: Mapped[str] = mapped_column(String(255)) + additional_context: Mapped[str] = mapped_column(Text, default="") + status: Mapped[ResearchRunStatus] = mapped_column( + enum_type(ResearchRunStatus, "research_run_status"), + default=ResearchRunStatus.QUEUED, + index=True, + ) + provider_response_id: Mapped[str | None] = mapped_column(String(255)) + input_snapshot: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + structured_result: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + markdown_report: Mapped[str | None] = mapped_column(Text) + error_message: Mapped[str | None] = mapped_column(Text) + started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + trend: Mapped[Trend | None] = relationship(back_populates="research_runs") + + +class ResearchSourceConfig(UUIDPrimaryKeyMixin, TimestampMixin, VersionMixin, Base): + __tablename__ = "research_source_configs" + __table_args__ = (UniqueConstraint("provider"),) + + provider: Mapped[str] = mapped_column(String(50), unique=True) + enabled: Mapped[bool] = mapped_column(Boolean, default=True) + config: Mapped[dict[str, Any]] = mapped_column(JsonType, default=dict) + + +class EtsyStatsImport(UUIDPrimaryKeyMixin, TimestampMixin, Base): + __tablename__ = "etsy_stats_imports" + __table_args__ = (UniqueConstraint("checksum"),) + + filename: Mapped[str] = mapped_column(String(255)) + checksum: Mapped[str] = mapped_column(String(64), unique=True, index=True) + period_start: Mapped[datetime] = mapped_column(DateTime(timezone=True)) + period_end: Mapped[datetime] = mapped_column(DateTime(timezone=True)) + row_count: Mapped[int] = mapped_column(Integer) + rows: Mapped[list[dict[str, Any]]] = mapped_column(JsonType, default=list) + warnings: Mapped[list[str]] = mapped_column(JsonType, default=list) class Product(UUIDPrimaryKeyMixin, TimestampMixin, VersionMixin, Base): @@ -209,6 +281,17 @@ class EtsyOAuthState(UUIDPrimaryKeyMixin, TimestampMixin, Base): consumed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) +class GoogleOAuthState(UUIDPrimaryKeyMixin, TimestampMixin, Base): + __tablename__ = "google_oauth_states" + + state_digest: Mapped[str] = mapped_column(String(64), unique=True, index=True) + code_verifier_ciphertext: Mapped[str] = mapped_column(Text) + redirect_uri: Mapped[str] = mapped_column(Text) + scopes: Mapped[list[str]] = mapped_column(JsonType, default=list) + expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + consumed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) + + class PostingSession(UUIDPrimaryKeyMixin, TimestampMixin, VersionMixin, Base): __tablename__ = "posting_sessions" @@ -298,6 +381,7 @@ class ImageGeneration(UUIDPrimaryKeyMixin, TimestampMixin, Base): parent_generation_id: Mapped[uuid.UUID | None] = mapped_column( ForeignKey("image_generations.id"), index=True ) + title: Mapped[str] = mapped_column(String(255), default="Untitled design") prompt: Mapped[str] = mapped_column(Text) status: Mapped[ImageGenerationStatus] = mapped_column( enum_type(ImageGenerationStatus, "image_generation_status"), @@ -416,8 +500,8 @@ class Job(UUIDPrimaryKeyMixin, TimestampMixin, Base): "job_type", "subject_id", unique=True, - postgresql_where=text("status IN ('queued', 'running')"), - sqlite_where=text("status IN ('queued', 'running')"), + postgresql_where=text("status IN ('queued', 'waiting', 'running')"), + sqlite_where=text("status IN ('queued', 'waiting', 'running')"), ), ) @@ -428,6 +512,7 @@ class Job(UUIDPrimaryKeyMixin, TimestampMixin, Base): enum_type(JobStatus, "job_status"), default=JobStatus.QUEUED, index=True ) attempts: Mapped[int] = mapped_column(Integer, default=0) + available_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), index=True) claimed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) error_message: Mapped[str | None] = mapped_column(Text) diff --git a/src/ecommerce_agent/domain/__init__.py b/src/ecommerce_agent/domain/__init__.py index 6b0e5d8..c92df14 100644 --- a/src/ecommerce_agent/domain/__init__.py +++ b/src/ecommerce_agent/domain/__init__.py @@ -5,6 +5,7 @@ JobStatus, ProductStatus, ProviderName, + ResearchRunStatus, TrendStatus, ) @@ -15,5 +16,6 @@ "JobStatus", "ProductStatus", "ProviderName", + "ResearchRunStatus", "TrendStatus", ] diff --git a/src/ecommerce_agent/domain/enums.py b/src/ecommerce_agent/domain/enums.py index bbd3ba9..eaa2919 100644 --- a/src/ecommerce_agent/domain/enums.py +++ b/src/ecommerce_agent/domain/enums.py @@ -2,11 +2,19 @@ class TrendStatus(StrEnum): + DISCOVERED = "discovered" PENDING_REVIEW = "pending_review" APPROVED = "approved" REJECTED = "rejected" +class ResearchRunStatus(StrEnum): + QUEUED = "queued" + RUNNING = "running" + COMPLETED = "completed" + FAILED = "failed" + + class ProductStatus(StrEnum): CONCEPT = "concept" GENERATING_DESIGN = "generating_design" @@ -37,6 +45,7 @@ class ApprovalDecision(StrEnum): class ProviderName(StrEnum): ETSY = "etsy" + GOOGLE_ANALYTICS = "google_analytics" PRINTFUL = "printful" OPENAI = "openai" MOCK = "mock" @@ -50,6 +59,7 @@ class ExternalOperationStatus(StrEnum): class JobStatus(StrEnum): QUEUED = "queued" + WAITING = "waiting" RUNNING = "running" SUCCEEDED = "succeeded" FAILED = "failed" diff --git a/src/ecommerce_agent/domain/research.py b/src/ecommerce_agent/domain/research.py new file mode 100644 index 0000000..f0ef954 --- /dev/null +++ b/src/ecommerce_agent/domain/research.py @@ -0,0 +1,99 @@ +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, HttpUrl, model_validator + + +class VerifiedSource(BaseModel): + model_config = ConfigDict(extra="forbid") + + title: str = Field(min_length=1, max_length=500) + url: HttpUrl + + +class TrendScoreBreakdown(BaseModel): + model_config = ConfigDict(extra="forbid") + + momentum: float = Field(ge=1, le=5) + audience_identity: float = Field(ge=1, le=5) + visual_suitability: float = Field(ge=1, le=5) + product_breadth: float = Field(ge=1, le=5) + printful_feasibility: float = Field(ge=1, le=5) + differentiation: float = Field(ge=1, le=5) + + +class DiscoveredTrend(BaseModel): + model_config = ConfigDict(extra="forbid") + + title: str = Field(min_length=1, max_length=255) + niche: str = Field(min_length=1, max_length=255) + blurb: str = Field(min_length=1, max_length=1200) + apparel_score: float = Field(ge=1, le=5) + score_breakdown: TrendScoreBreakdown + confidence: float = Field(ge=0, le=1) + evidence: list[str] = Field(min_length=1, max_length=10) + verified_sources: list[VerifiedSource] = Field(min_length=1, max_length=12) + risk_warnings: list[str] = Field(max_length=10) + + +class TrendDiscoveryResult(BaseModel): + model_config = ConfigDict(extra="forbid") + + summary: str = Field(min_length=1, max_length=2000) + trends: list[DiscoveredTrend] = Field(min_length=1, max_length=25) + + +class ProductProposal(BaseModel): + model_config = ConfigDict(extra="forbid") + + title: str = Field(min_length=1, max_length=255) + garment_type: str = Field(min_length=1, max_length=100) + printful_catalog_product_id: int | None + printful_product_title: str | None = Field(max_length=255) + technique: str = Field(min_length=1, max_length=100) + placement: str = Field(min_length=1, max_length=100) + creative_concept: str = Field(min_length=1, max_length=2000) + visual_direction: str = Field(min_length=1, max_length=2000) + audience: str = Field(min_length=1, max_length=1000) + rationale: str = Field(min_length=1, max_length=2000) + differentiation: str = Field(min_length=1, max_length=2000) + production_complexity: Literal["low", "medium", "high"] + + +class ResearchReport(BaseModel): + model_config = ConfigDict(extra="forbid") + + executive_recommendation: str = Field(min_length=1, max_length=4000) + apparel_score: float = Field(ge=1, le=5) + confidence: float = Field(ge=0, le=1) + trend_definition_and_timing: str = Field(min_length=1, max_length=8000) + evidence_and_source_signals: str = Field(min_length=1, max_length=10000) + audience_communities_and_buying_occasions: str = Field(min_length=1, max_length=8000) + apparel_and_printful_suitability: str = Field(min_length=1, max_length=8000) + proposals: list[ProductProposal] = Field(min_length=1, max_length=10) + competition_and_differentiation: str = Field(min_length=1, max_length=8000) + risk_warnings: list[str] = Field(max_length=20) + keywords_and_validation_experiments: list[str] = Field(min_length=1, max_length=30) + unknowns_and_next_actions: list[str] = Field(min_length=1, max_length=30) + verified_sources: list[VerifiedSource] = Field(min_length=1, max_length=40) + + @model_validator(mode="after") + def validate_printful_matches(self) -> "ResearchReport": + for proposal in self.proposals: + if (proposal.printful_catalog_product_id is None) != ( + proposal.printful_product_title is None + ): + raise ValueError( + "Printful catalog product ID and product title must be supplied together." + ) + return self + + +class ResearchTaskResult(BaseModel): + model_config = ConfigDict(extra="forbid") + + status: Literal["queued", "in_progress", "completed", "retryable", "failed"] + response_id: str | None = None + result: dict[str, Any] | None = None + error_message: str | None = None + retry_after_seconds: int | None = Field(default=None, ge=1, le=3600) + restart_required: bool = False diff --git a/src/ecommerce_agent/jobs/handlers.py b/src/ecommerce_agent/jobs/handlers.py index 1057417..1ff011d 100644 --- a/src/ecommerce_agent/jobs/handlers.py +++ b/src/ecommerce_agent/jobs/handlers.py @@ -4,6 +4,10 @@ from ecommerce_agent.db.models import PrintfulProductDraft from ecommerce_agent.pipeline import steps +from ecommerce_agent.pipeline.research_reports import ( + run_research_report, + run_trend_discovery, +) from ecommerce_agent.services import ServiceContainer @@ -17,4 +21,9 @@ async def submit_printful_draft( return await submit_draft(session, services, draft_id) -JOB_HANDLERS = {**steps.STEPS, "submit_printful_draft": submit_printful_draft} +JOB_HANDLERS = { + **steps.STEPS, + "submit_printful_draft": submit_printful_draft, + "discover_trends": run_trend_discovery, + "generate_research_report": run_research_report, +} diff --git a/src/ecommerce_agent/jobs/queue.py b/src/ecommerce_agent/jobs/queue.py index 58be348..9458d45 100644 --- a/src/ecommerce_agent/jobs/queue.py +++ b/src/ecommerce_agent/jobs/queue.py @@ -7,7 +7,7 @@ from ecommerce_agent.db.models import Job from ecommerce_agent.domain.enums import JobStatus -ACTIVE_STATUSES = (JobStatus.QUEUED, JobStatus.RUNNING) +ACTIVE_STATUSES = (JobStatus.QUEUED, JobStatus.WAITING, JobStatus.RUNNING) async def enqueue( diff --git a/src/ecommerce_agent/jobs/worker.py b/src/ecommerce_agent/jobs/worker.py index 0733e25..87ac4ef 100644 --- a/src/ecommerce_agent/jobs/worker.py +++ b/src/ecommerce_agent/jobs/worker.py @@ -9,7 +9,7 @@ from ecommerce_agent.db.models import Job from ecommerce_agent.domain.enums import JobStatus from ecommerce_agent.jobs.handlers import JOB_HANDLERS -from ecommerce_agent.pipeline.errors import ConflictError +from ecommerce_agent.pipeline.errors import ConflictError, DeferJob from ecommerce_agent.services import ServiceContainer logger = logging.getLogger(__name__) @@ -32,6 +32,10 @@ async def claim_next(session: AsyncSession) -> Job | None: .where( or_( Job.status == JobStatus.QUEUED, + and_( + Job.status == JobStatus.WAITING, + or_(Job.available_at.is_(None), Job.available_at <= now), + ), and_(Job.status == JobStatus.RUNNING, Job.claimed_at < stale_before), ) ) @@ -45,6 +49,7 @@ async def claim_next(session: AsyncSession) -> Job | None: return None job.status = JobStatus.RUNNING job.claimed_at = now + job.available_at = None job.attempts += 1 await session.commit() return job @@ -64,6 +69,17 @@ async def process_one( if handler is None: raise ValueError(f"Unknown job type: {job_type}") await handler(session, services, subject_id) + except DeferJob as deferred: + await session.rollback() + job = await session.get(Job, job_id) + if job is not None: + job.status = JobStatus.WAITING + job.claimed_at = None + job.available_at = datetime.now(UTC) + timedelta(seconds=deferred.seconds) + job.attempts = max(0, job.attempts - 1) + job.error_message = str(deferred) + await session.commit() + return True except ConflictError as error: # The subject is no longer in the state this step requires, which # means the step's outcome was already committed before a crash. diff --git a/src/ecommerce_agent/pipeline/errors.py b/src/ecommerce_agent/pipeline/errors.py index 0d7cfb8..db33d93 100644 --- a/src/ecommerce_agent/pipeline/errors.py +++ b/src/ecommerce_agent/pipeline/errors.py @@ -4,3 +4,9 @@ class ConflictError(RuntimeError): class NotFoundError(RuntimeError): pass + + +class DeferJob(Exception): + def __init__(self, *, seconds: int, message: str = "Waiting for external work."): + super().__init__(message) + self.seconds = seconds diff --git a/src/ecommerce_agent/pipeline/research_reports.py b/src/ecommerce_agent/pipeline/research_reports.py new file mode 100644 index 0000000..fcc4d94 --- /dev/null +++ b/src/ecommerce_agent/pipeline/research_reports.py @@ -0,0 +1,440 @@ +import uuid +from datetime import UTC, datetime +from decimal import Decimal +from typing import Any, cast + +from sqlalchemy import desc, select +from sqlalchemy.ext.asyncio import AsyncSession + +from ecommerce_agent.db.models import ( + EtsyStatsImport, + PrintfulCatalogProduct, + Product, + ResearchRun, + Trend, + TrendDiscoveryRun, +) +from ecommerce_agent.domain.enums import ResearchRunStatus, TrendStatus +from ecommerce_agent.domain.research import ( + ResearchReport, + ResearchTaskResult, + TrendDiscoveryResult, +) +from ecommerce_agent.pipeline.errors import ConflictError, DeferJob +from ecommerce_agent.services import ServiceContainer +from ecommerce_agent.services.google_analytics import ( + GoogleAnalyticsError, + GoogleTokenCipher, + collect_ga4_context, +) + +MAX_PROVIDER_RETRIES = 3 + + +async def run_trend_discovery( + session: AsyncSession, services: ServiceContainer, run_id: uuid.UUID +) -> TrendDiscoveryRun: + run = await session.get(TrendDiscoveryRun, run_id) + if run is None: + raise ConflictError("Trend discovery run not found.") + if run.status == ResearchRunStatus.COMPLETED: + raise ConflictError("Trend discovery run is already complete.") + try: + if run.provider_response_id: + task = await services.research.retrieve( + run.provider_response_id, kind="trend_discovery" + ) + else: + run.status = ResearchRunStatus.RUNNING + run.started_at = run.started_at or datetime.now(UTC) + if run.input_snapshot: + context = _provider_context(run.input_snapshot) + else: + context = await collect_research_context(session, services) + run.input_snapshot = context + await session.commit() + task = await services.research.start_trend_discovery( + context=context, + limit=services.research_trend_limit, + ) + run.provider_response_id = task.response_id or run.provider_response_id + if task.status in {"queued", "in_progress"}: + await session.commit() + raise DeferJob(seconds=services.research_poll_seconds) + if task.status == "retryable": + await _defer_provider_retry( + session, run, task, poll_seconds=services.research_poll_seconds + ) + if task.status == "failed" or task.result is None: + raise RuntimeError(task.error_message or "Trend discovery failed.") + result = TrendDiscoveryResult.model_validate(task.result) + existing = ( + await session.scalars(select(Trend).where(Trend.discovery_run_id == run.id)) + ).all() + if not existing: + for position, item in enumerate(result.trends): + session.add( + Trend( + discovery_run_id=run.id, + research_batch_id=run.id, + position=position, + title=item.title, + niche=item.niche, + source="openai-web" if task.response_id else "mock-research", + evidence=item.evidence, + research={}, + blurb=item.blurb, + apparel_score=Decimal(str(item.apparel_score)), + score_breakdown=item.score_breakdown.model_dump(mode="json"), + confidence=Decimal(str(item.confidence)), + verified_sources=[ + source.model_dump(mode="json") for source in item.verified_sources + ], + risk_warnings=item.risk_warnings, + status=TrendStatus.DISCOVERED, + ) + ) + run.structured_result = result.model_dump(mode="json") + run.input_snapshot = _provider_context(run.input_snapshot) + run.status = ResearchRunStatus.COMPLETED + run.finished_at = datetime.now(UTC) + run.error_message = None + await session.commit() + return run + except DeferJob: + raise + except Exception as error: + await session.rollback() + run = await session.get(TrendDiscoveryRun, run_id) + if run is not None: + run.status = ResearchRunStatus.FAILED + run.error_message = str(error) + run.finished_at = datetime.now(UTC) + await session.commit() + return cast(TrendDiscoveryRun, run) + + +async def run_research_report( + session: AsyncSession, services: ServiceContainer, run_id: uuid.UUID +) -> ResearchRun: + run = await session.get(ResearchRun, run_id) + if run is None: + raise ConflictError("Research run not found.") + if run.status == ResearchRunStatus.COMPLETED: + raise ConflictError("Research run is already complete.") + try: + if run.provider_response_id: + task = await services.research.retrieve( + run.provider_response_id, kind="research_report" + ) + else: + run.status = ResearchRunStatus.RUNNING + run.started_at = run.started_at or datetime.now(UTC) + if run.input_snapshot: + context = _provider_context(run.input_snapshot) + else: + context = await collect_research_context(session, services) + if run.trend_id: + trend = await session.get(Trend, run.trend_id) + if trend is not None: + context["selected_trend"] = { + "title": trend.title, + "niche": trend.niche, + "blurb": trend.blurb, + "apparel_score": _number(trend.apparel_score), + "confidence": _number(trend.confidence), + "evidence": trend.evidence, + "sources": trend.verified_sources, + "risk_warnings": trend.risk_warnings, + } + run.input_snapshot = context + await session.commit() + task = await services.research.start_research_report( + topic_title=run.topic_title, + additional_context=run.additional_context, + context=context, + ) + run.provider_response_id = task.response_id or run.provider_response_id + if task.status in {"queued", "in_progress"}: + await session.commit() + raise DeferJob(seconds=services.research_poll_seconds) + if task.status == "retryable": + await _defer_provider_retry( + session, run, task, poll_seconds=services.research_poll_seconds + ) + if task.status == "failed" or task.result is None: + if await _restart_manual_retry(session, run, task): + raise DeferJob( + seconds=services.research_poll_seconds, + message="Restarting the failed OpenAI research response.", + ) + raise RuntimeError(task.error_message or "Research report failed.") + report = ResearchReport.model_validate(task.result) + run.structured_result = report.model_dump(mode="json") + run.markdown_report = render_markdown(run, report) + run.input_snapshot = _provider_context(run.input_snapshot) + run.status = ResearchRunStatus.COMPLETED + run.finished_at = datetime.now(UTC) + run.error_message = None + await session.commit() + return run + except DeferJob: + raise + except Exception as error: + await session.rollback() + run = await session.get(ResearchRun, run_id) + if run is not None: + run.status = ResearchRunStatus.FAILED + run.error_message = str(error) + run.finished_at = datetime.now(UTC) + await session.commit() + return cast(ResearchRun, run) + + +async def _defer_provider_retry( + session: AsyncSession, + run: TrendDiscoveryRun | ResearchRun, + task: ResearchTaskResult, + *, + poll_seconds: int, +) -> None: + snapshot = dict(run.input_snapshot) + retry_count = int(snapshot.get("_workshop_provider_retries") or 0) + 1 + if retry_count > MAX_PROVIDER_RETRIES: + raise RuntimeError( + task.error_message + or f"OpenAI research exceeded {MAX_PROVIDER_RETRIES} provider retries." + ) + snapshot["_workshop_provider_retries"] = retry_count + run.input_snapshot = snapshot + if task.restart_required: + run.provider_response_id = None + run.error_message = task.error_message + await session.commit() + retry_seconds = max(task.retry_after_seconds or poll_seconds, poll_seconds) + raise DeferJob( + seconds=retry_seconds, + message=( + f"OpenAI research is temporarily unavailable; retrying " + f"({retry_count}/{MAX_PROVIDER_RETRIES})." + ), + ) + + +def _provider_context(snapshot: dict[str, Any]) -> dict[str, Any]: + return {key: value for key, value in snapshot.items() if not key.startswith("_workshop_")} + + +async def _restart_manual_retry( + session: AsyncSession, + run: ResearchRun, + task: ResearchTaskResult, +) -> bool: + snapshot = dict(run.input_snapshot) + if not snapshot.pop("_workshop_manual_retry", False) or not run.provider_response_id: + return False + run.input_snapshot = snapshot + run.provider_response_id = None + run.error_message = task.error_message + await session.commit() + return True + + +async def collect_research_context( + session: AsyncSession, services: ServiceContainer +) -> dict[str, Any]: + etsy_stats = ( + await session.scalars( + select(EtsyStatsImport).order_by(desc(EtsyStatsImport.created_at)).limit(5) + ) + ).all() + printful_products = ( + await session.scalars( + select(PrintfulCatalogProduct) + .where( + PrintfulCatalogProduct.us_available.is_(True), + PrintfulCatalogProduct.discontinued.is_(False), + ) + .order_by( + PrintfulCatalogProduct.favorite.desc(), + PrintfulCatalogProduct.workshop_score.desc(), + ) + .limit(30) + ) + ).all() + products = ( + await session.scalars( + select(Product) + .where(Product.etsy_listing_id.is_not(None)) + .order_by(desc(Product.updated_at)) + .limit(20) + ) + ).all() + etsy_signals: list[dict[str, Any]] = [] + etsy_warnings: list[str] = [] + for product in products: + if not product.etsy_listing_id: + continue + try: + listing = await services.etsy.get_listing(product.etsy_listing_id) + transactions = await services.etsy.get_transactions(product.etsy_listing_id) + etsy_signals.append( + { + "listing_id": product.etsy_listing_id, + "title": product.title, + "views": listing.get("views"), + "favorites": listing.get("num_favorers"), + "transactions": transactions[:50], + } + ) + except Exception as error: + etsy_warnings.append(f"Could not collect listing {product.etsy_listing_id}: {error}") + ga4: dict[str, Any] + try: + cipher = ( + GoogleTokenCipher(services.google_token_encryption_key) + if services.google_token_encryption_key is not None + else None + ) + ga4 = await collect_ga4_context( + session, + cipher=cipher, + client_id=services.google_client_id, + client_secret=services.google_client_secret, + ) + except GoogleAnalyticsError as error: + ga4 = {"connected": True, "rows": [], "warning": str(error)} + return { + "trend_limit": services.research_trend_limit, + "etsy_stats_imports": [ + { + "period_start": item.period_start.isoformat(), + "period_end": item.period_end.isoformat(), + "rows": sorted( + item.rows, + key=lambda row: int(row.get("visits") or 0), + reverse=True, + )[:200], + } + for item in etsy_stats + ], + "etsy_listing_signals": etsy_signals, + "etsy_warnings": etsy_warnings, + "google_analytics": ga4, + "google_trends": { + "enabled": False, + "reason": "Official API access remains limited alpha.", + }, + "printful_catalog": [ + { + "catalog_product_id": item.catalog_product_id, + "title": item.title, + "category_id": item.category_id, + "favorite": item.favorite, + "manual_rating": item.manual_rating, + "workshop_score": _number(item.workshop_score), + "cheapest_us_price": _number(item.cheapest_us_price), + "preferred_placement": item.preferred_placement, + "techniques": item.raw_payload.get("techniques", []), + } + for item in printful_products + ], + } + + +def render_markdown(run: ResearchRun, report: ResearchReport) -> str: + completed = run.finished_at or datetime.now(UTC) + lines = [ + f"# {run.topic_title}", + "", + "## Run Metadata", + "", + f"- Research run: `{run.id}`", + f"- Completed: {completed.isoformat()}", + f"- Source trend: `{run.trend_id}`" if run.trend_id else "- Source trend: Manual topic", + "", + "## Executive Recommendation", + "", + report.executive_recommendation, + "", + f"- Apparel recommendation: **{report.apparel_score:.1f}/5**", + f"- Confidence: **{report.confidence:.0%}**", + "", + "## Trend Definition And Timing", + "", + report.trend_definition_and_timing, + "", + "## Evidence And Source Signals", + "", + report.evidence_and_source_signals, + "", + "## Audience, Communities, And Buying Occasions", + "", + report.audience_communities_and_buying_occasions, + "", + "## Apparel And Printful Suitability", + "", + report.apparel_and_printful_suitability, + "", + "## Product And Artwork Proposals", + "", + ] + for index, proposal in enumerate(report.proposals, start=1): + lines.extend( + [ + f"### {index}. {proposal.title}", + "", + f"- Garment: {proposal.garment_type}", + ( + f"- Printful match: `{proposal.printful_catalog_product_id}` " + f"{proposal.printful_product_title}" + if proposal.printful_catalog_product_id is not None + else "- Printful match: No exact cached product selected" + ), + f"- Technique and placement: {proposal.technique}, {proposal.placement}", + f"- Audience: {proposal.audience}", + f"- Production complexity: {proposal.production_complexity}", + "", + f"**Creative concept:** {proposal.creative_concept}", + "", + f"**Visual direction:** {proposal.visual_direction}", + "", + f"**Rationale:** {proposal.rationale}", + "", + f"**Differentiation:** {proposal.differentiation}", + "", + ] + ) + lines.extend( + [ + "## Competition And Differentiation", + "", + report.competition_and_differentiation, + "", + "## Trademark, Copyright, Cultural, And Etsy-Policy Warnings", + "", + *(_bullets(report.risk_warnings) or ["- No specific warnings identified."]), + "", + "## Keywords And Validation Experiments", + "", + *_bullets(report.keywords_and_validation_experiments), + "", + "## Unknowns And Recommended Next Actions", + "", + *_bullets(report.unknowns_and_next_actions), + "", + "## Verified Sources", + "", + *[f"- [{source.title}]({source.url})" for source in report.verified_sources], + "", + ] + ) + return "\n".join(lines) + + +def _bullets(values: list[str]) -> list[str]: + return [f"- {value}" for value in values] + + +def _number(value: Decimal | None) -> float | None: + return float(value) if value is not None else None diff --git a/src/ecommerce_agent/services/factory.py b/src/ecommerce_agent/services/factory.py index 7239e25..a0d57e9 100644 --- a/src/ecommerce_agent/services/factory.py +++ b/src/ecommerce_agent/services/factory.py @@ -15,6 +15,7 @@ ImageGenerationProvider, PrintfulGateway, PublicAssetStore, + ResearchProvider, TrendSource, ) from ecommerce_agent.services.mocks import ( @@ -27,6 +28,7 @@ from ecommerce_agent.services.openai_copy import OpenAICopyProvider from ecommerce_agent.services.openai_images import OpenAIImageProvider from ecommerce_agent.services.printful import DisabledPrintfulGateway, PrintfulClient +from ecommerce_agent.services.research import MockResearchProvider, OpenAIResearchProvider @dataclass(frozen=True) @@ -35,12 +37,18 @@ class ServiceContainer: image: ImageGenerationProvider image_assets: LocalAssetStore trends: TrendSource + research: ResearchProvider etsy: EtsyGateway printful: PrintfulGateway printful_etsy: PrintfulGateway print_assets: PublicAssetStore | None print_asset_url_ttl_seconds: int print_asset_mode: str | None + research_poll_seconds: int + research_trend_limit: int + google_client_id: str | None + google_client_secret: str | None + google_token_encryption_key: str | None def build_service_container(settings: Settings | None = None) -> ServiceContainer: @@ -65,9 +73,15 @@ def build_service_container(settings: Settings | None = None) -> ServiceContaine api_key=settings.openai_api_key.get_secret_value(), model=settings.openai_text_model, ) + research_provider: ResearchProvider = OpenAIResearchProvider( + api_key=settings.openai_api_key.get_secret_value(), + model=settings.openai_text_model, + max_tool_calls=settings.research_max_tool_calls, + ) else: image = MockImageGenerationProvider(assets) copy = MockCopyGenerationProvider() + research_provider = MockResearchProvider() if settings.printful_mode == "real": if not settings.integrations_enabled: @@ -143,10 +157,24 @@ def build_service_container(settings: Settings | None = None) -> ServiceContaine image=image, image_assets=assets, trends=MockTrendSource(), + research=research_provider, etsy=etsy, printful=printful, printful_etsy=printful_etsy, print_assets=print_assets, print_asset_url_ttl_seconds=settings.printful_asset_url_ttl_seconds, print_asset_mode=print_asset_mode, + research_poll_seconds=settings.research_poll_seconds, + research_trend_limit=settings.research_trend_limit, + google_client_id=settings.google_client_id, + google_client_secret=( + settings.google_client_secret.get_secret_value() + if settings.google_client_secret is not None + else None + ), + google_token_encryption_key=( + settings.google_token_encryption_key.get_secret_value() + if settings.google_token_encryption_key is not None + else None + ), ) diff --git a/src/ecommerce_agent/services/google_analytics.py b/src/ecommerce_agent/services/google_analytics.py new file mode 100644 index 0000000..096804a --- /dev/null +++ b/src/ecommerce_agent/services/google_analytics.py @@ -0,0 +1,388 @@ +import base64 +import hashlib +import secrets +from collections.abc import Callable +from datetime import UTC, datetime, timedelta +from typing import Any, cast +from urllib.parse import urlencode + +import httpx +from cryptography.fernet import Fernet, InvalidToken +from sqlalchemy import delete, select +from sqlalchemy.ext.asyncio import AsyncSession + +from ecommerce_agent.db.models import ( + GoogleOAuthState, + OAuthCredential, + ResearchSourceConfig, +) +from ecommerce_agent.domain.enums import ProviderName + +GOOGLE_AUTHORIZATION_URL = "https://accounts.google.com/o/oauth2/v2/auth" +GOOGLE_TOKEN_URL = "https://oauth2.googleapis.com/token" +GOOGLE_USERINFO_URL = "https://openidconnect.googleapis.com/v1/userinfo" +GOOGLE_ANALYTICS_ADMIN_URL = "https://analyticsadmin.googleapis.com/v1beta/accountSummaries" +GOOGLE_ANALYTICS_DATA_URL = "https://analyticsdata.googleapis.com/v1beta" +GOOGLE_ANALYTICS_SCOPES = ( + "openid", + "email", + "https://www.googleapis.com/auth/analytics.readonly", +) +ACCESS_TOKEN_SKEW = timedelta(minutes=5) + + +class GoogleAnalyticsError(RuntimeError): + pass + + +class GoogleTokenCipher: + def __init__(self, key: str): + try: + self._fernet = Fernet(key.encode()) + except (TypeError, ValueError) as error: + raise GoogleAnalyticsError( + "GOOGLE_TOKEN_ENCRYPTION_KEY must be a valid Fernet key." + ) from error + + def encrypt(self, value: str) -> str: + return self._fernet.encrypt(value.encode()).decode() + + def decrypt(self, value: str) -> str: + try: + return self._fernet.decrypt(value.encode()).decode() + except InvalidToken as error: + raise GoogleAnalyticsError( + "Stored Google credentials cannot be decrypted with the configured key." + ) from error + + +async def create_google_authorization_url( + session: AsyncSession, + *, + cipher: GoogleTokenCipher, + client_id: str, + redirect_uri: str, + ttl_seconds: int, + clock: Callable[[], datetime] | None = None, +) -> str: + now = (clock or (lambda: datetime.now(UTC)))() + state = secrets.token_urlsafe(32) + verifier = secrets.token_urlsafe(64) + challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).decode() + challenge = challenge.rstrip("=") + session.add( + GoogleOAuthState( + state_digest=_digest(state), + code_verifier_ciphertext=cipher.encrypt(verifier), + redirect_uri=redirect_uri, + scopes=list(GOOGLE_ANALYTICS_SCOPES), + expires_at=now + timedelta(seconds=ttl_seconds), + ) + ) + await session.flush() + parameters = { + "client_id": client_id, + "redirect_uri": redirect_uri, + "response_type": "code", + "scope": " ".join(GOOGLE_ANALYTICS_SCOPES), + "access_type": "offline", + "include_granted_scopes": "true", + "prompt": "consent", + "state": state, + "code_challenge": challenge, + "code_challenge_method": "S256", + } + return f"{GOOGLE_AUTHORIZATION_URL}?{urlencode(parameters)}" + + +async def complete_google_authorization( + session: AsyncSession, + *, + cipher: GoogleTokenCipher, + client_id: str, + client_secret: str, + state: str, + code: str, + client: httpx.AsyncClient, + clock: Callable[[], datetime] | None = None, +) -> OAuthCredential: + now = (clock or (lambda: datetime.now(UTC)))() + oauth_state = await session.scalar( + select(GoogleOAuthState) + .where(GoogleOAuthState.state_digest == _digest(state)) + .with_for_update() + ) + if oauth_state is None: + raise GoogleAnalyticsError("Google OAuth state is invalid.") + if oauth_state.consumed_at is not None: + raise GoogleAnalyticsError("Google OAuth state was already used.") + if _as_utc(oauth_state.expires_at) <= _as_utc(now): + raise GoogleAnalyticsError("Google OAuth state expired.") + try: + response = await client.post( + GOOGLE_TOKEN_URL, + data={ + "client_id": client_id, + "client_secret": client_secret, + "code": code, + "code_verifier": cipher.decrypt(oauth_state.code_verifier_ciphertext), + "grant_type": "authorization_code", + "redirect_uri": oauth_state.redirect_uri, + }, + ) + response.raise_for_status() + token = cast(dict[str, Any], response.json()) + access_token = str(token["access_token"]) + refresh_token = str(token["refresh_token"]) + expires_in = int(token["expires_in"]) + except (httpx.HTTPError, KeyError, TypeError, ValueError) as error: + raise GoogleAnalyticsError("Google OAuth token exchange failed.") from error + try: + user_response = await client.get( + GOOGLE_USERINFO_URL, + headers={"Authorization": f"Bearer {access_token}"}, + ) + user_response.raise_for_status() + user = cast(dict[str, Any], user_response.json()) + email = str(user["email"]) + except (httpx.HTTPError, KeyError, TypeError, ValueError) as error: + raise GoogleAnalyticsError("Google account lookup failed.") from error + + credential = await session.scalar( + select(OAuthCredential).where( + OAuthCredential.provider == ProviderName.GOOGLE_ANALYTICS, + OAuthCredential.account_id == email, + ) + ) + if credential is None: + credential = OAuthCredential( + provider=ProviderName.GOOGLE_ANALYTICS, + account_id=email, + refresh_token_ciphertext=cipher.encrypt(refresh_token), + ) + session.add(credential) + credential.access_token_ciphertext = cipher.encrypt(access_token) + credential.refresh_token_ciphertext = cipher.encrypt(refresh_token) + credential.token_type = str(token.get("token_type", "Bearer")) + credential.scopes = str(token.get("scope", "")).split() + credential.user_id = str(user.get("sub") or "") + credential.access_token_expires_at = now + timedelta(seconds=expires_in) + credential.refresh_token_expires_at = None + oauth_state.consumed_at = now + await session.flush() + return credential + + +async def google_access_token( + session: AsyncSession, + *, + cipher: GoogleTokenCipher, + client_id: str, + client_secret: str, + client: httpx.AsyncClient, +) -> tuple[OAuthCredential, str]: + credential = await session.scalar( + select(OAuthCredential) + .where(OAuthCredential.provider == ProviderName.GOOGLE_ANALYTICS) + .order_by(OAuthCredential.updated_at.desc()) + .with_for_update() + ) + if credential is None: + raise GoogleAnalyticsError("Google Analytics is not connected.") + now = datetime.now(UTC) + if ( + credential.access_token_ciphertext + and credential.access_token_expires_at + and _as_utc(credential.access_token_expires_at) > now + ACCESS_TOKEN_SKEW + ): + return credential, cipher.decrypt(credential.access_token_ciphertext) + try: + response = await client.post( + GOOGLE_TOKEN_URL, + data={ + "client_id": client_id, + "client_secret": client_secret, + "refresh_token": cipher.decrypt(credential.refresh_token_ciphertext), + "grant_type": "refresh_token", + }, + ) + response.raise_for_status() + token = cast(dict[str, Any], response.json()) + access_token = str(token["access_token"]) + expires_in = int(token["expires_in"]) + except (httpx.HTTPError, KeyError, TypeError, ValueError) as error: + raise GoogleAnalyticsError("Google OAuth token refresh failed.") from error + credential.access_token_ciphertext = cipher.encrypt(access_token) + credential.access_token_expires_at = now + timedelta(seconds=expires_in) + credential.token_type = str(token.get("token_type", credential.token_type)) + if token.get("scope"): + credential.scopes = str(token["scope"]).split() + await session.flush() + return credential, access_token + + +async def list_ga4_properties( + session: AsyncSession, + *, + cipher: GoogleTokenCipher, + client_id: str, + client_secret: str, + client: httpx.AsyncClient, +) -> list[dict[str, str]]: + _, token = await google_access_token( + session, + cipher=cipher, + client_id=client_id, + client_secret=client_secret, + client=client, + ) + try: + response = await client.get( + GOOGLE_ANALYTICS_ADMIN_URL, + params={"pageSize": 200}, + headers={"Authorization": f"Bearer {token}"}, + ) + response.raise_for_status() + payload = cast(dict[str, Any], response.json()) + except (httpx.HTTPError, TypeError, ValueError) as error: + raise GoogleAnalyticsError("Google Analytics property lookup failed.") from error + properties: list[dict[str, str]] = [] + for account in cast(list[dict[str, Any]], payload.get("accountSummaries", [])): + for prop in cast(list[dict[str, Any]], account.get("propertySummaries", [])): + name = str(prop.get("property") or "") + if name.startswith("properties/"): + properties.append( + { + "property_id": name.removeprefix("properties/"), + "display_name": str(prop.get("displayName") or name), + "account_name": str(account.get("displayName") or ""), + } + ) + return properties + + +async def select_ga4_property( + session: AsyncSession, *, property_id: str, display_name: str +) -> ResearchSourceConfig: + config = await session.scalar( + select(ResearchSourceConfig).where(ResearchSourceConfig.provider == "google_analytics") + ) + if config is None: + config = ResearchSourceConfig(provider="google_analytics", enabled=True) + session.add(config) + config.enabled = True + config.config = {"property_id": property_id, "display_name": display_name} + await session.flush() + return config + + +async def collect_ga4_context( + session: AsyncSession, + *, + cipher: GoogleTokenCipher | None, + client_id: str | None, + client_secret: str | None, + client: httpx.AsyncClient | None = None, +) -> dict[str, Any]: + config = await session.scalar( + select(ResearchSourceConfig).where(ResearchSourceConfig.provider == "google_analytics") + ) + if ( + config is None + or not config.enabled + or not config.config.get("property_id") + or cipher is None + or not client_id + or not client_secret + ): + return {"connected": False, "rows": []} + owns_client = client is None + http_client = client or httpx.AsyncClient(timeout=30) + try: + _, token = await google_access_token( + session, + cipher=cipher, + client_id=client_id, + client_secret=client_secret, + client=http_client, + ) + property_id = str(config.config["property_id"]) + response = await http_client.post( + f"{GOOGLE_ANALYTICS_DATA_URL}/properties/{property_id}:runReport", + headers={"Authorization": f"Bearer {token}"}, + json={ + "dateRanges": [ + {"startDate": "28daysAgo", "endDate": "yesterday", "name": "recent"}, + {"startDate": "56daysAgo", "endDate": "29daysAgo", "name": "prior"}, + ], + "dimensions": [ + {"name": "landingPagePlusQueryString"}, + {"name": "pageTitle"}, + {"name": "sessionSourceMedium"}, + ], + "metrics": [ + {"name": "sessions"}, + {"name": "engagedSessions"}, + {"name": "keyEvents"}, + {"name": "totalRevenue"}, + ], + "limit": "100", + }, + ) + response.raise_for_status() + payload = cast(dict[str, Any], response.json()) + except (httpx.HTTPError, TypeError, ValueError) as error: + raise GoogleAnalyticsError("Google Analytics report collection failed.") from error + finally: + if owns_client: + await http_client.aclose() + dimension_names = [ + str(item.get("name")) + for item in cast(list[dict[str, Any]], payload.get("dimensionHeaders", [])) + ] + metric_names = [ + str(item.get("name")) + for item in cast(list[dict[str, Any]], payload.get("metricHeaders", [])) + ] + rows = [] + for row in cast(list[dict[str, Any]], payload.get("rows", [])): + dimensions = cast(list[dict[str, Any]], row.get("dimensionValues", [])) + metrics = cast(list[dict[str, Any]], row.get("metricValues", [])) + rows.append( + { + **{ + name: str(value.get("value", "")) + for name, value in zip(dimension_names, dimensions, strict=False) + }, + **{ + name: str(value.get("value", "")) + for name, value in zip(metric_names, metrics, strict=False) + }, + } + ) + return { + "connected": True, + "property_id": config.config["property_id"], + "display_name": config.config.get("display_name"), + "rows": rows, + } + + +async def disconnect_google_analytics(session: AsyncSession) -> None: + await session.execute( + delete(OAuthCredential).where(OAuthCredential.provider == ProviderName.GOOGLE_ANALYTICS) + ) + config = await session.scalar( + select(ResearchSourceConfig).where(ResearchSourceConfig.provider == "google_analytics") + ) + if config is not None: + config.enabled = False + config.config = {} + + +def _digest(value: str) -> str: + return hashlib.sha256(value.encode()).hexdigest() + + +def _as_utc(value: datetime) -> datetime: + return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC) diff --git a/src/ecommerce_agent/services/interfaces.py b/src/ecommerce_agent/services/interfaces.py index 0a4fdaa..ac81079 100644 --- a/src/ecommerce_agent/services/interfaces.py +++ b/src/ecommerce_agent/services/interfaces.py @@ -9,6 +9,7 @@ ImageReference, TrendCandidate, ) +from ecommerce_agent.domain.research import ResearchTaskResult class ImageGenerationProvider(Protocol): @@ -35,6 +36,22 @@ class TrendSource(Protocol): async def discover(self, limit: int = 5) -> list[TrendCandidate]: ... +class ResearchProvider(Protocol): + async def start_trend_discovery( + self, *, context: Mapping[str, Any], limit: int + ) -> ResearchTaskResult: ... + + async def start_research_report( + self, + *, + topic_title: str, + additional_context: str, + context: Mapping[str, Any], + ) -> ResearchTaskResult: ... + + async def retrieve(self, response_id: str, *, kind: str) -> ResearchTaskResult: ... + + class ListingCopy(BaseModel): model_config = ConfigDict(extra="forbid") diff --git a/src/ecommerce_agent/services/openai_images.py b/src/ecommerce_agent/services/openai_images.py index 6c29bdc..37b3905 100644 --- a/src/ecommerce_agent/services/openai_images.py +++ b/src/ecommerce_agent/services/openai_images.py @@ -47,11 +47,14 @@ def estimate_image_output_cost(*, model: str, quality: str, width: int, height: if model != "gpt-image-2" and not model.startswith("gpt-image-2-"): raise OpenAIImageError(f"No configured cost estimate is available for model {model}.") estimate = _OUTPUT_COSTS_USD.get((quality, width, height)) - if estimate is None: - raise OpenAIImageError( - f"No configured cost estimate is available for {quality} {width}x{height}." - ) - return estimate + if estimate is not None: + return estimate + square_estimate = _OUTPUT_COSTS_USD.get((quality, 1024, 1024)) + if square_estimate is None: + raise OpenAIImageError(f"No configured cost estimate is available for {quality}.") + return (square_estimate * Decimal(width * height) / Decimal(1024 * 1024)).quantize( + Decimal("0.001") + ) def build_effective_prompt(request: ImageGenerationRequest) -> str: diff --git a/src/ecommerce_agent/services/research.py b/src/ecommerce_agent/services/research.py new file mode 100644 index 0000000..e25d6d7 --- /dev/null +++ b/src/ecommerce_agent/services/research.py @@ -0,0 +1,468 @@ +import json +import re +from collections.abc import Mapping +from typing import Any, Literal, cast +from urllib.parse import urlsplit, urlunsplit + +import httpx + +from ecommerce_agent.domain.research import ( + DiscoveredTrend, + ProductProposal, + ResearchReport, + ResearchTaskResult, + TrendDiscoveryResult, + TrendScoreBreakdown, + VerifiedSource, +) + +OPENAI_RESPONSES_URL = "https://api.openai.com/v1/responses" +DEFAULT_RETRY_SECONDS = 30 +RETRYABLE_ERROR_CODES = {"rate_limit_exceeded", "server_error"} + + +class ResearchProviderError(RuntimeError): + pass + + +class MockResearchProvider: + async def start_trend_discovery( + self, *, context: Mapping[str, Any], limit: int + ) -> ResearchTaskResult: + del context + source = VerifiedSource.model_validate( + { + "title": "Etsy Shop Stats guidance", + "url": ( + "https://help.etsy.com/hc/en-us/articles/" + "115015774268-How-to-Use-Etsy-Stats-for-Your-Shop" + ), + } + ) + seeds = [ + ( + "Bookish Garden Club", + "reading and gardening crossover", + "A cozy identity trend combining reading rituals, botanical motifs, and " + "small-club language that translates cleanly to understated apparel.", + ), + ( + "Retro Pickleball Weekends", + "recreational pickleball", + "Social pickleball culture continues to support playful, giftable apparel " + "with club, weekend, and tournament-adjacent themes.", + ), + ( + "Cozy Astronomy Society", + "casual stargazing", + "Accessible astronomy, night-sky outings, and society-style graphics create " + "an evergreen visual niche for simple garments.", + ), + ] + trends = [ + DiscoveredTrend( + title=title, + niche=niche, + blurb=blurb, + apparel_score=4.1 - index * 0.2, + score_breakdown=TrendScoreBreakdown( + momentum=4, + audience_identity=4.5, + visual_suitability=4.5, + product_breadth=4, + printful_feasibility=4.5, + differentiation=3.5, + ), + confidence=0.72, + evidence=[ + "Identity-oriented language supports wearable self-expression.", + "The visual vocabulary works for both print and embroidery.", + ], + verified_sources=[source], + risk_warnings=[], + ) + for index, (title, niche, blurb) in enumerate(seeds[:limit]) + ] + result = TrendDiscoveryResult( + summary="Deterministic mock trend discovery for local development.", + trends=trends, + ) + return ResearchTaskResult(status="completed", result=result.model_dump(mode="json")) + + async def start_research_report( + self, + *, + topic_title: str, + additional_context: str, + context: Mapping[str, Any], + ) -> ResearchTaskResult: + products = cast(list[dict[str, Any]], context.get("printful_catalog", [])) + matched = products[0] if products else None + source = VerifiedSource.model_validate( + { + "title": "Etsy Shop Stats guidance", + "url": ( + "https://help.etsy.com/hc/en-us/articles/" + "115015774268-How-to-Use-Etsy-Stats-for-Your-Shop" + ), + } + ) + proposal = ProductProposal( + title=f"{topic_title} Society Cap", + garment_type="hat", + printful_catalog_product_id=( + int(matched["catalog_product_id"]) if matched is not None else None + ), + printful_product_title=str(matched["title"]) if matched is not None else None, + technique="embroidery", + placement="front", + creative_concept=f"A restrained society-style emblem for {topic_title}.", + visual_direction="Simple line work, two thread colors, compact badge silhouette.", + audience="Gift buyers and people who identify with the topic.", + rationale="The compact motif is legible, wearable, and suitable for repeat variants.", + differentiation="Avoid generic slogans by using a specific visual ritual or symbol.", + production_complexity="low", + ) + report = ResearchReport( + executive_recommendation=( + f"Test {topic_title} with one focused embroidered hat and one printed shirt. " + f"Operator context: {additional_context or 'none supplied'}." + ), + apparel_score=4.2, + confidence=0.7, + trend_definition_and_timing=( + f"{topic_title} is treated as an identity and gift-oriented apparel theme." + ), + evidence_and_source_signals="Mock evidence is used in local development mode.", + audience_communities_and_buying_occasions=( + "Core buyers include enthusiasts, club members, and occasion-based gift buyers." + ), + apparel_and_printful_suitability=( + "Compact marks suit embroidery; larger illustrations suit shirts and hoodies." + ), + proposals=[proposal], + competition_and_differentiation=( + "Differentiate through restrained visual systems and specific subculture details." + ), + risk_warnings=["Run a trademark and copyright check before publishing."], + keywords_and_validation_experiments=[ + f"{topic_title.lower()} hat", + f"{topic_title.lower()} shirt", + "Compare embroidered badge and typography-led artwork.", + ], + unknowns_and_next_actions=[ + "Validate current Etsy search language.", + "Review the first mockup at actual garment scale.", + ], + verified_sources=[source], + ) + return ResearchTaskResult(status="completed", result=report.model_dump(mode="json")) + + async def retrieve(self, response_id: str, *, kind: str) -> ResearchTaskResult: + del response_id, kind + return ResearchTaskResult( + status="failed", + error_message="Mock research tasks complete synchronously.", + ) + + +class OpenAIResearchProvider: + def __init__( + self, + *, + api_key: str, + model: str, + max_tool_calls: int, + client: httpx.AsyncClient | None = None, + ): + self._api_key = api_key + self._model = model + self._max_tool_calls = max_tool_calls + self._client = client + + async def start_trend_discovery( + self, *, context: Mapping[str, Any], limit: int + ) -> ResearchTaskResult: + prompt = ( + "Discover emerging US-facing trends relevant to a small Etsy apparel shop selling " + "simple hats, shirts, hoodies, prints, and embroidery. Focus on pop culture, " + "clothing, fashion, hobbies, aesthetics, communities, seasonal behavior, and " + "giftable identity signals. Do not merely list today's largest news stories. " + f"Return at most {limit} distinct trends. Score apparel fit independently from IP, " + "trademark, copyright, cultural, or marketplace-policy risk; put those concerns in " + "risk_warnings. Every verified_sources URL must be a URL actually consulted by web " + "search. Prefer recent, primary, or credible market sources. Use the supplied " + f"first-party and catalog context as supporting signals:\n{json.dumps(dict(context))}" + ) + return await self._start( + prompt=prompt, + schema_name="trend_discovery", + schema=_openai_json_schema(TrendDiscoveryResult.model_json_schema()), + reasoning_effort="medium", + ) + + async def start_research_report( + self, + *, + topic_title: str, + additional_context: str, + context: Mapping[str, Any], + ) -> ResearchTaskResult: + prompt = ( + "Create an evidence-backed apparel opportunity report for an Etsy and Printful " + "operator. Investigate the topic deeply, including what it means, why it may matter " + "now, audiences, communities, buying occasions, competing apparel, visual language, " + "and practical product directions. Produce 1 to 10 proposals. When supplied Printful " + "catalog products genuinely fit, cite their exact catalog_product_id and title; " + "otherwise leave both fields null. Do not create designs or products. Keep commercial " + "scoring independent from legal and policy warnings. Every verified_sources URL must " + "be a URL actually consulted by web search.\n\n" + f"Topic title: {topic_title}\n" + f"Additional operator context: {additional_context or 'None supplied.'}\n\n" + f"First-party and Printful context:\n{json.dumps(dict(context))}" + ) + return await self._start( + prompt=prompt, + schema_name="research_report", + schema=_openai_json_schema(ResearchReport.model_json_schema()), + reasoning_effort="high", + ) + + async def retrieve(self, response_id: str, *, kind: str) -> ResearchTaskResult: + owns_client = self._client is None + client = self._client or httpx.AsyncClient(timeout=httpx.Timeout(90)) + try: + response = await client.get( + f"{OPENAI_RESPONSES_URL}/{response_id}", + headers={"Authorization": f"Bearer {self._api_key}"}, + params=[("include[]", "web_search_call.action.sources")], + ) + if _is_retryable_http(response): + return _http_retry_result(response, response_id=response_id) + if response.is_error: + raise ResearchProviderError( + "OpenAI research retrieval failed with HTTP " + f"{response.status_code}: {_error_detail(response)}" + ) + return _parse_response(response, kind=kind) + finally: + if owns_client: + await client.aclose() + + async def _start( + self, + *, + prompt: str, + schema_name: str, + schema: dict[str, Any], + reasoning_effort: str, + ) -> ResearchTaskResult: + payload = { + "model": self._model, + "background": True, + "store": True, + "reasoning": {"effort": reasoning_effort}, + "tools": [{"type": "web_search", "search_context_size": "high"}], + "tool_choice": "required", + "max_tool_calls": self._max_tool_calls, + "include": ["web_search_call.action.sources"], + "input": prompt, + "text": { + "format": { + "type": "json_schema", + "name": schema_name, + "strict": True, + "schema": schema, + } + }, + } + owns_client = self._client is None + client = self._client or httpx.AsyncClient(timeout=httpx.Timeout(90)) + try: + response = await client.post( + OPENAI_RESPONSES_URL, + headers={"Authorization": f"Bearer {self._api_key}"}, + json=payload, + ) + if _is_retryable_http(response): + return _http_retry_result(response, response_id=None) + if response.is_error: + raise ResearchProviderError( + "OpenAI research start failed with HTTP " + f"{response.status_code}: {_error_detail(response)}" + ) + return _parse_response(response, kind=schema_name) + finally: + if owns_client: + await client.aclose() + + +def _parse_response(response: httpx.Response, *, kind: str) -> ResearchTaskResult: + try: + payload = cast(dict[str, Any], response.json()) + status = cast( + Literal["queued", "in_progress", "completed", "failed"], + str(payload["status"]), + ) + response_id = str(payload["id"]) + except (KeyError, TypeError, ValueError) as error: + raise ResearchProviderError("OpenAI returned an invalid research response.") from error + if status in {"queued", "in_progress"}: + return ResearchTaskResult(status=status, response_id=response_id) + if status != "completed": + error_detail = payload.get("error") or payload.get("incomplete_details") or status + if _provider_error_code(error_detail) in RETRYABLE_ERROR_CODES: + error_message = f"OpenAI research ended with {error_detail}." + return ResearchTaskResult( + status="retryable", + response_id=response_id, + error_message=error_message, + retry_after_seconds=_retry_after_seconds(error_message), + restart_required=True, + ) + return ResearchTaskResult( + status="failed", + response_id=response_id, + error_message=f"OpenAI research ended with {error_detail}.", + ) + + output_text = _output_text(payload) + consulted_urls = _consulted_urls(payload) + try: + parsed = json.loads(output_text) + if kind in {"trend_discovery", "discovery"}: + discovery_result = TrendDiscoveryResult.model_validate(parsed) + verified_trends = [] + for trend in discovery_result.trends: + verified_sources = _verified_sources(trend.verified_sources, consulted_urls) + if verified_sources: + verified_trends.append( + trend.model_copy(update={"verified_sources": verified_sources}) + ) + if not verified_trends: + raise ResearchProviderError( + "OpenAI returned no trends with verified web-search sources." + ) + discovery_result = discovery_result.model_copy(update={"trends": verified_trends}) + result_payload = discovery_result.model_dump(mode="json") + elif kind in {"research_report", "report"}: + report_result = ResearchReport.model_validate(parsed) + verified_sources = _verified_sources(report_result.verified_sources, consulted_urls) + if not verified_sources: + raise ResearchProviderError( + "OpenAI returned no verified web-search sources for the report." + ) + report_result = report_result.model_copy(update={"verified_sources": verified_sources}) + result_payload = report_result.model_dump(mode="json") + else: + raise ResearchProviderError(f"Unknown research result kind: {kind}.") + except (TypeError, ValueError) as error: + raise ResearchProviderError("OpenAI returned invalid structured research.") from error + return ResearchTaskResult( + status="completed", + response_id=response_id, + result=result_payload, + ) + + +def _openai_json_schema(value: Any) -> Any: + if isinstance(value, dict): + return {key: _openai_json_schema(item) for key, item in value.items() if key != "format"} + if isinstance(value, list): + return [_openai_json_schema(item) for item in value] + return value + + +def _output_text(payload: Mapping[str, Any]) -> str: + for output in cast(list[dict[str, Any]], payload.get("output", [])): + if output.get("type") != "message": + continue + for content in cast(list[dict[str, Any]], output.get("content", [])): + if content.get("type") == "output_text" and content.get("text"): + return str(content["text"]) + if content.get("type") == "refusal": + raise ResearchProviderError("OpenAI refused the research request.") + raise ResearchProviderError("OpenAI returned no structured research output.") + + +def _consulted_urls(payload: Mapping[str, Any]) -> set[str]: + urls: set[str] = set() + for output in cast(list[dict[str, Any]], payload.get("output", [])): + if output.get("type") == "web_search_call": + action = cast(dict[str, Any], output.get("action", {})) + for source in cast(list[dict[str, Any]], action.get("sources", [])): + if source.get("url"): + urls.add(str(source["url"])) + if output.get("type") == "message": + for content in cast(list[dict[str, Any]], output.get("content", [])): + for annotation in cast(list[dict[str, Any]], content.get("annotations", [])): + if annotation.get("type") == "url_citation" and annotation.get("url"): + urls.add(str(annotation["url"])) + return urls + + +def _error_detail(response: httpx.Response) -> str: + try: + payload = cast(dict[str, Any], response.json()) + error = cast(dict[str, Any], payload.get("error", {})) + message = str(error.get("message") or "").strip() + if message: + return message[:1000] + except (TypeError, ValueError): + pass + return "no provider detail" + + +def _verified_sources( + sources: list[VerifiedSource], consulted_urls: set[str] +) -> list[VerifiedSource]: + consulted = {_canonical_url(url) for url in consulted_urls} + return [source for source in sources if _canonical_url(str(source.url)) in consulted] + + +def _canonical_url(value: str) -> str: + parsed = urlsplit(value) + hostname = (parsed.hostname or "").lower() + port = parsed.port + if port and not ( + (parsed.scheme.lower() == "http" and port == 80) + or (parsed.scheme.lower() == "https" and port == 443) + ): + hostname = f"{hostname}:{port}" + path = parsed.path.rstrip("/") or "/" + return urlunsplit((parsed.scheme.lower(), hostname, path, parsed.query, "")) + + +def _is_retryable_http(response: httpx.Response) -> bool: + return response.status_code in {408, 409, 429} or response.status_code >= 500 + + +def _http_retry_result(response: httpx.Response, *, response_id: str | None) -> ResearchTaskResult: + detail = _error_detail(response) + return ResearchTaskResult( + status="retryable", + response_id=response_id, + error_message=( + f"OpenAI research request temporarily failed with HTTP {response.status_code}: {detail}" + ), + retry_after_seconds=_retry_after_seconds(response.headers.get("retry-after") or detail), + restart_required=response_id is None, + ) + + +def _provider_error_code(error_detail: Any) -> str: + if isinstance(error_detail, Mapping): + return str(error_detail.get("code") or "") + return "" + + +def _retry_after_seconds(value: str) -> int: + header_value = value.strip() + try: + return max(1, min(3600, int(float(header_value)))) + except ValueError: + pass + match = re.search(r"try again in\s+(\d+(?:\.\d+)?)s", value, re.IGNORECASE) + if match: + return max(1, min(3600, int(float(match.group(1))) + 1)) + return DEFAULT_RETRY_SECONDS diff --git a/tests/integration/test_postgres_schema.py b/tests/integration/test_postgres_schema.py index b8ceb61..bd6bffe 100644 --- a/tests/integration/test_postgres_schema.py +++ b/tests/integration/test_postgres_schema.py @@ -76,9 +76,15 @@ async def schema_details() -> tuple[set[str], set[str], set[str], set[str]]: "printful_categories", "printful_catalog_products", "printful_product_drafts", + "trend_discovery_runs", + "research_runs", + "research_source_configs", + "etsy_stats_imports", + "google_oauth_states", }.issubset(tables) assert { "parent_generation_id", + "title", "output_format", "background_removal_friendly", }.issubset(image_generation_columns) diff --git a/tests/smoke/test_api_and_dashboard.py b/tests/smoke/test_api_and_dashboard.py index d1ffebb..69ec095 100644 --- a/tests/smoke/test_api_and_dashboard.py +++ b/tests/smoke/test_api_and_dashboard.py @@ -2,17 +2,23 @@ import io import uuid from dataclasses import replace +from datetime import UTC, datetime, timedelta from urllib.parse import urlsplit from fastapi.testclient import TestClient from httpx import ASGITransport, AsyncClient from PIL import Image +from sqlalchemy import select 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 from ecommerce_agent.jobs.worker import process_one from ecommerce_agent.services.assets import SignedLocalAssetStore +from ecommerce_agent.services.research import MockResearchProvider def make_settings(tmp_path) -> Settings: @@ -42,6 +48,263 @@ def test_health_endpoint(tmp_path) -> None: assert response.json() == {"status": "ok"} +async def test_trend_discovery_and_research_report_api(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: + discovery = await client.post("/api/trend-discovery-runs") + assert discovery.status_code == 201 + assert await process_one(db_sessions, services) + + trends = (await client.get("/api/trends")).json() + discovered = [item for item in trends if item["status"] == "discovered"] + assert len(discovered) == 3 + assert discovered[0]["blurb"] + assert 1 <= float(discovered[0]["apparel_score"]) <= 5 + + created = await client.post( + "/api/research-runs", + json={ + "topic_title": discovered[0]["title"], + "additional_context": "Prioritize embroidery.", + "trend_id": discovered[0]["id"], + }, + ) + assert created.status_code == 201 + research_id = created.json()["id"] + assert await process_one(db_sessions, services) + + report = (await client.get(f"/api/research-runs/{research_id}")).json() + assert report["status"] == "completed" + assert "## Product And Artwork Proposals" in report["markdown_report"] + download = await client.get(f"/api/research-runs/{research_id}/download") + assert download.status_code == 200 + assert download.headers["content-type"].startswith("text/markdown") + assert "attachment;" in download.headers["content-disposition"] + + +async def test_etsy_stats_import_validation_and_deduplication( + 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 + content = ( + b"period_start,period_end,search_term,visits,listing_id,listing_title,notes\n" + b"2026-05-01,2026-05-31,bookish hat,12,123,Book Hat,\n" + b"2026-05-01,2026-05-31,garden club shirt,7,,,seasonal\n" + ) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + template = await client.get("/api/research-sources/etsy-stats-imports/template") + assert template.status_code == 200 + assert "search_term" in template.text + imported = await client.post( + "/api/research-sources/etsy-stats-imports", + files={"upload": ("etsy.csv", content, "text/csv")}, + ) + assert imported.status_code == 201 + assert imported.json()["row_count"] == 2 + duplicate = await client.post( + "/api/research-sources/etsy-stats-imports", + files={"upload": ("etsy.csv", content, "text/csv")}, + ) + assert duplicate.status_code == 409 + history = (await client.get("/api/research-sources/etsy-stats-imports")).json() + assert history[0]["row_count"] == 2 + + +async def test_background_research_polling_does_not_consume_attempts(db_sessions, services) -> None: + mock = MockResearchProvider() + + class DeferredResearchProvider: + async def start_trend_discovery(self, *, context, limit): + del context, limit + raise AssertionError("Unexpected discovery call") + + async def start_research_report(self, *, topic_title, additional_context, context): + del topic_title, additional_context, context + from ecommerce_agent.domain.research import ResearchTaskResult + + return ResearchTaskResult(status="queued", response_id="resp_deferred") + + async def retrieve(self, response_id, *, kind): + assert response_id == "resp_deferred" + assert kind == "research_report" + return await mock.start_research_report( + topic_title="Deferred topic", + additional_context="", + context={}, + ) + + configured = replace( + services, + research=DeferredResearchProvider(), + research_poll_seconds=1, + ) + async with db_sessions() as session: + run = ResearchRun( + topic_title="Deferred topic", + additional_context="", + status=ResearchRunStatus.QUEUED, + ) + session.add(run) + await session.flush() + from ecommerce_agent.jobs.queue import enqueue + + await enqueue(session, "generate_research_report", run.id) + await session.commit() + run_id = run.id + + assert await process_one(db_sessions, configured) + async with db_sessions() as session: + job = await session.scalar(select(Job).where(Job.subject_id == run_id)) + assert job is not None + assert job.status == JobStatus.WAITING + assert job.attempts == 0 + job.available_at = datetime.now(UTC) - timedelta(seconds=1) + await session.commit() + + assert await process_one(db_sessions, configured) + async with db_sessions() as session: + run = await session.get(ResearchRun, run_id) + assert run is not None + assert run.status == ResearchRunStatus.COMPLETED + + +async def test_background_research_restarts_after_transient_rate_limit( + db_sessions, services +) -> None: + mock = MockResearchProvider() + + class RateLimitedResearchProvider: + starts = 0 + + async def start_trend_discovery(self, *, context, limit): + del context, limit + raise AssertionError("Unexpected discovery call") + + async def start_research_report(self, *, topic_title, additional_context, context): + del topic_title, additional_context, context + from ecommerce_agent.domain.research import ResearchTaskResult + + self.starts += 1 + if self.starts == 1: + return ResearchTaskResult(status="queued", response_id="resp_rate_limited") + return await mock.start_research_report( + topic_title="Rate-limited topic", + additional_context="", + context={}, + ) + + async def retrieve(self, response_id, *, kind): + assert response_id == "resp_rate_limited" + assert kind == "research_report" + from ecommerce_agent.domain.research import ResearchTaskResult + + return ResearchTaskResult( + status="retryable", + response_id=response_id, + error_message="Temporary OpenAI rate limit.", + retry_after_seconds=1, + restart_required=True, + ) + + provider = RateLimitedResearchProvider() + configured = replace( + services, + research=provider, + research_poll_seconds=1, + ) + async with db_sessions() as session: + run = ResearchRun( + topic_title="Rate-limited topic", + additional_context="", + status=ResearchRunStatus.QUEUED, + ) + session.add(run) + await session.flush() + from ecommerce_agent.jobs.queue import enqueue + + await enqueue(session, "generate_research_report", run.id) + await session.commit() + run_id = run.id + + for expected_starts in (1, 1): + assert await process_one(db_sessions, configured) + async with db_sessions() as session: + job = await session.scalar(select(Job).where(Job.subject_id == run_id)) + assert job is not None + assert job.status == JobStatus.WAITING + assert job.attempts == 0 + assert provider.starts == expected_starts + job.available_at = datetime.now(UTC) - timedelta(seconds=1) + await session.commit() + + async with db_sessions() as session: + run = await session.get(ResearchRun, run_id) + assert run is not None + assert run.provider_response_id is None + assert run.input_snapshot["_workshop_provider_retries"] == 1 + + assert await process_one(db_sessions, configured) + async with db_sessions() as session: + run = await session.get(ResearchRun, run_id) + assert run is not None + assert run.status == ResearchRunStatus.COMPLETED + assert "_workshop_provider_retries" not in run.input_snapshot + assert provider.starts == 2 + + +async def test_research_retry_reuses_completed_provider_response( + 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 db_sessions() as session: + run = ResearchRun( + topic_title="Reusable response", + additional_context="", + status=ResearchRunStatus.FAILED, + provider_response_id="resp_completed", + input_snapshot={"catalog": []}, + error_message="Citation validation failed.", + ) + session.add(run) + await session.commit() + run_id = run.id + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post(f"/api/research-runs/{run_id}/retry") + assert response.status_code == 200 + + async with db_sessions() as session: + run = await session.get(ResearchRun, run_id) + assert run is not None + assert run.status == ResearchRunStatus.QUEUED + assert run.provider_response_id == "resp_completed" + assert run.input_snapshot["_workshop_manual_retry"] is True + job = await session.scalar(select(Job).where(Job.subject_id == run_id)) + assert job is not None + assert job.status == JobStatus.QUEUED + + async def test_standalone_image_generation_history_and_asset_serving( tmp_path, db_sessions, services ) -> None: @@ -56,18 +319,27 @@ async def _get_session(): async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: response = await client.post( "/api/image-generations", - json={"prompt": "A clean geometric hat patch concept"}, + json={ + "title": "Hat patch", + "prompt": "A clean geometric hat patch concept", + "width": 1024, + "height": 1280, + }, ) assert response.status_code == 201 generation = response.json() assert generation["status"] == "succeeded" + assert generation["title"] == "Hat patch" assert generation["asset_uri"].startswith("/api/assets/") + assert generation["width"] == 1024 + assert generation["height"] == 1280 assert generation["estimated_cost_usd"] == 0.0 assert generation["output_format"] == "png" assert generation["background_removal_friendly"] is True history = (await client.get("/api/image-generations")).json() assert len(history) == 1 + assert history[0]["title"] == "Hat patch" assert history[0]["prompt"] == "A clean geometric hat patch concept" asset = await client.get(generation["asset_uri"]) @@ -78,6 +350,8 @@ async def _get_session(): f"/api/image-generations/{generation['id']}/edits", data={ "prompt": "Make the shape blue", + "width": "1536", + "height": "768", "output_format": "webp", "background_removal_friendly": "false", }, @@ -92,11 +366,21 @@ async def _get_session(): edited = edit_response.json() assert edited["status"] == "succeeded" assert edited["parent_generation_id"] == generation["id"] + assert edited["title"] == "Hat patch" + assert edited["width"] == 1536 + assert edited["height"] == 768 assert edited["output_format"] == "webp" assert edited["background_removal_friendly"] is False assert edited["generation_metadata"]["reference_count"] == 2 assert edited["asset_uri"].endswith(".webp") + renamed = await client.patch( + f"/api/image-generations/{edited['id']}", + json={"title": "Blue hat patch"}, + ) + assert renamed.status_code == 200 + assert renamed.json()["title"] == "Blue hat patch" + delete_parent = await client.delete(f"/api/image-generations/{generation['id']}") assert delete_parent.status_code == 204 assert (await client.get(generation["asset_uri"])).status_code == 404 @@ -151,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_config.py b/tests/unit/test_config.py index 2b9d26b..e4ad686 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -11,9 +11,19 @@ def test_real_provider_requires_top_level_kill_switch() -> None: with pytest.raises(ValidationError, match="INTEGRATIONS_ENABLED"): + Settings( + _env_file=None, + openai_mode="real", + openai_api_key="test-key", + ) + + +def test_test_env_rejects_real_openai() -> None: + with pytest.raises(ValidationError, match="OPENAI_MODE=real"): Settings( _env_file=None, app_env="test", + integrations_enabled=True, openai_mode="real", openai_api_key="test-key", ) @@ -23,7 +33,6 @@ def test_real_openai_requires_api_key() -> None: with pytest.raises(ValidationError, match="OPENAI_API_KEY"): Settings( _env_file=None, - app_env="test", integrations_enabled=True, openai_mode="real", ) @@ -52,13 +61,11 @@ def test_factory_uses_mock_gateways(tmp_path) -> None: assert isinstance(services.printful, MockPrintfulGateway) -def test_factory_can_use_real_openai_with_mock_commerce_gateways(tmp_path) -> None: +def test_factory_can_use_real_openai_with_mock_etsy_and_disabled_printful(tmp_path) -> None: settings = Settings( _env_file=None, - app_env="test", asset_root=tmp_path, integrations_enabled=True, - printful_mode="mock", openai_mode="real", openai_api_key="test-key", ) @@ -66,7 +73,7 @@ def test_factory_can_use_real_openai_with_mock_commerce_gateways(tmp_path) -> No assert isinstance(services.image, OpenAIImageProvider) assert isinstance(services.etsy, MockEtsyGateway) - assert isinstance(services.printful, MockPrintfulGateway) + assert isinstance(services.printful, DisabledPrintfulGateway) def test_factory_disables_printful_without_using_mock_data(tmp_path) -> None: 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" diff --git a/tests/unit/test_google_analytics.py b/tests/unit/test_google_analytics.py new file mode 100644 index 0000000..c4f8ff9 --- /dev/null +++ b/tests/unit/test_google_analytics.py @@ -0,0 +1,94 @@ +import httpx +from cryptography.fernet import Fernet + +from ecommerce_agent.services.google_analytics import ( + GoogleTokenCipher, + complete_google_authorization, + create_google_authorization_url, + list_ga4_properties, +) + + +async def test_google_oauth_and_property_discovery(db_sessions) -> None: + cipher = GoogleTokenCipher(Fernet.generate_key().decode()) + async with db_sessions() as session: + url = await create_google_authorization_url( + session, + cipher=cipher, + client_id="google-client", + redirect_uri="http://localhost/callback", + ttl_seconds=600, + ) + await session.commit() + state = url.split("state=", 1)[1].split("&", 1)[0] + + async def handler(request: httpx.Request) -> httpx.Response: + if request.url.path == "/token": + return httpx.Response( + 200, + request=request, + json={ + "access_token": "access", + "refresh_token": "refresh", + "expires_in": 3600, + "token_type": "Bearer", + "scope": "openid email https://www.googleapis.com/auth/analytics.readonly", + }, + ) + if request.url.path == "/v1/userinfo": + return httpx.Response( + 200, + request=request, + json={"email": "owner@example.com", "sub": "google-user"}, + ) + if request.url.path == "/v1beta/accountSummaries": + return httpx.Response( + 200, + request=request, + json={ + "accountSummaries": [ + { + "displayName": "Workshop", + "propertySummaries": [ + { + "property": "properties/12345", + "displayName": "WearKR", + } + ], + } + ] + }, + ) + return httpx.Response(404, request=request) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + async with db_sessions() as session: + credential = await complete_google_authorization( + session, + cipher=cipher, + client_id="google-client", + client_secret="google-secret", + state=state, + code="authorization-code", + client=client, + ) + await session.commit() + assert credential.account_id == "owner@example.com" + assert cipher.decrypt(credential.refresh_token_ciphertext) == "refresh" + + async with db_sessions() as session: + properties = await list_ga4_properties( + session, + cipher=cipher, + client_id="google-client", + client_secret="google-secret", + client=client, + ) + await session.commit() + assert properties == [ + { + "property_id": "12345", + "display_name": "WearKR", + "account_name": "Workshop", + } + ] diff --git a/tests/unit/test_openai_images.py b/tests/unit/test_openai_images.py index c127ae3..51de8a4 100644 --- a/tests/unit/test_openai_images.py +++ b/tests/unit/test_openai_images.py @@ -16,6 +16,7 @@ OpenAIImageError, OpenAIImageOutcomeUnknown, OpenAIImageProvider, + estimate_image_output_cost, ) @@ -60,6 +61,12 @@ def provider( ) +def test_estimates_custom_size_by_pixels() -> None: + assert estimate_image_output_cost( + model="gpt-image-2", quality="medium", width=1024, height=1280 + ) == Decimal("0.066") + + async def test_generates_png_with_expected_request_and_metadata(tmp_path) -> None: captured: dict[str, Any] = {} diff --git a/tests/unit/test_research.py b/tests/unit/test_research.py new file mode 100644 index 0000000..588f7d3 --- /dev/null +++ b/tests/unit/test_research.py @@ -0,0 +1,255 @@ +import json +import uuid +from datetime import UTC, datetime + +import httpx +import pytest + +from ecommerce_agent.db.models import ResearchRun +from ecommerce_agent.pipeline.research_reports import render_markdown +from ecommerce_agent.services.research import ( + MockResearchProvider, + OpenAIResearchProvider, + ResearchProviderError, +) + + +async def test_openai_research_accepts_only_consulted_sources() -> None: + mock = MockResearchProvider() + completed = await mock.start_trend_discovery(context={}, limit=1) + result = completed.result + assert result is not None + source_url = result["trends"][0]["verified_sources"][0]["url"] + + async def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v1/responses" + request_payload = json.loads(request.content) + source_schema = request_payload["text"]["format"]["schema"]["$defs"]["VerifiedSource"] + assert "format" not in source_schema["properties"]["url"] + return httpx.Response( + 200, + request=request, + json={ + "id": "resp_research", + "status": "completed", + "output": [ + { + "type": "web_search_call", + "action": { + "type": "search", + "sources": [{"url": source_url, "title": "Etsy"}], + }, + }, + { + "type": "message", + "content": [ + { + "type": "output_text", + "text": json.dumps(result), + "annotations": [], + } + ], + }, + ], + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAIResearchProvider( + api_key="test", + model="gpt-test", + max_tool_calls=4, + client=client, + ) + task = await provider.start_trend_discovery(context={}, limit=1) + assert task.status == "completed" + assert task.response_id == "resp_research" + + +async def test_openai_research_rejects_results_without_verified_source_urls() -> None: + mock = MockResearchProvider() + completed = await mock.start_trend_discovery(context={}, limit=1) + result = completed.result + assert result is not None + + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + request=request, + json={ + "id": "resp_bad_source", + "status": "completed", + "output": [ + { + "type": "web_search_call", + "action": { + "type": "search", + "sources": [{"url": "https://example.com/other"}], + }, + }, + { + "type": "message", + "content": [ + { + "type": "output_text", + "text": json.dumps(result), + "annotations": [], + } + ], + }, + ], + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAIResearchProvider( + api_key="test", + model="gpt-test", + max_tool_calls=4, + client=client, + ) + with pytest.raises(ResearchProviderError, match="no trends with verified"): + await provider.start_trend_discovery(context={}, limit=1) + + +async def test_openai_research_discards_only_unverified_source_urls() -> None: + mock = MockResearchProvider() + completed = await mock.start_trend_discovery(context={}, limit=1) + result = completed.result + assert result is not None + verified_url = result["trends"][0]["verified_sources"][0]["url"] + result["trends"][0]["verified_sources"].append( + {"title": "Unsupported source", "url": "https://example.com/not-consulted"} + ) + + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + request=request, + json={ + "id": "resp_partial_sources", + "status": "completed", + "output": [ + { + "type": "web_search_call", + "action": { + "type": "search", + "sources": [{"url": verified_url, "title": "Etsy"}], + }, + }, + { + "type": "message", + "content": [ + { + "type": "output_text", + "text": json.dumps(result), + "annotations": [], + } + ], + }, + ], + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAIResearchProvider( + api_key="test", + model="gpt-test", + max_tool_calls=4, + client=client, + ) + task = await provider.start_trend_discovery(context={}, limit=1) + assert task.status == "completed" + assert task.result is not None + assert task.result["trends"][0]["verified_sources"] == [ + {"title": "Etsy Shop Stats guidance", "url": verified_url} + ] + + +async def test_openai_research_retrieval_includes_web_search_sources() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v1/responses/resp_pending" + assert request.url.params.get_list("include[]") == ["web_search_call.action.sources"] + return httpx.Response( + 200, + request=request, + json={"id": "resp_pending", "status": "in_progress", "output": []}, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAIResearchProvider( + api_key="test", + model="gpt-test", + max_tool_calls=4, + client=client, + ) + task = await provider.retrieve("resp_pending", kind="trend_discovery") + assert task.status == "in_progress" + assert task.response_id == "resp_pending" + + +async def test_openai_research_marks_background_rate_limits_retryable() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + request=request, + json={ + "id": "resp_rate_limited", + "status": "failed", + "error": { + "code": "rate_limit_exceeded", + "message": "Please try again in 1.43s.", + }, + "output": [], + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAIResearchProvider( + api_key="test", + model="gpt-test", + max_tool_calls=4, + client=client, + ) + task = await provider.retrieve("resp_rate_limited", kind="trend_discovery") + assert task.status == "retryable" + assert task.retry_after_seconds == 2 + assert task.restart_required is True + + +async def test_markdown_report_has_stable_outline() -> None: + mock = MockResearchProvider() + completed = await mock.start_research_report( + topic_title="Bookish Garden Club", + additional_context="Focus on embroidery.", + context={}, + ) + assert completed.result is not None + from ecommerce_agent.domain.research import ResearchReport + + report = ResearchReport.model_validate(completed.result) + run = ResearchRun( + id=uuid.uuid4(), + topic_title="Bookish Garden Club", + additional_context="Focus on embroidery.", + finished_at=datetime(2026, 6, 15, tzinfo=UTC), + ) + markdown = render_markdown(run, report) + expected_headings = [ + "# Bookish Garden Club", + "## Run Metadata", + "## Executive Recommendation", + "## Trend Definition And Timing", + "## Evidence And Source Signals", + "## Audience, Communities, And Buying Occasions", + "## Apparel And Printful Suitability", + "## Product And Artwork Proposals", + "## Competition And Differentiation", + "## Trademark, Copyright, Cultural, And Etsy-Policy Warnings", + "## Keywords And Validation Experiments", + "## Unknowns And Recommended Next Actions", + "## Verified Sources", + ] + positions = [markdown.index(heading) for heading in expected_headings] + assert positions == sorted(positions) + assert "Apparel recommendation: **4.2/5**" in markdown