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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 24 additions & 11 deletions apps/api/app/api/v1/routes/jobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

import os
import uuid
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Literal, Optional, cast
from urllib.parse import urlparse

Expand Down Expand Up @@ -53,6 +53,7 @@
StandardErrorObject,
)
from shared.services.storage.file_upload_service import FileUploadService
from shared.utils.utc_now import utc_now_naive
from shared.utils.url_security import (
validate_http_url_and_resolve_ip_async,
)
Expand Down Expand Up @@ -268,6 +269,15 @@ def ensure_utc(dt: Optional[datetime]) -> Optional[datetime]:
return dt.replace(tzinfo=timezone.utc)


def normalize_naive_utc_filter_datetime(dt: Optional[datetime]) -> Optional[datetime]:
"""Convert a query datetime into the naive UTC form used by database columns."""
if dt is None:
return None
if dt.tzinfo is None or dt.utcoffset() is None:
return dt
return dt.astimezone(timezone.utc).replace(tzinfo=None)


def require_utc(dt: Optional[datetime], *, field_name: str) -> datetime:
"""Normalize a required datetime to UTC."""
normalized_dt = ensure_utc(dt)
Expand Down Expand Up @@ -678,23 +688,28 @@ async def list_jobs(
user_message="recent_days only supports 1, 7, or 30",
violations=[{"field": "recent_days", "description": "Invalid value"}],
)
created_after = None
created_after: Optional[datetime] = None
if recent_days:
from datetime import datetime, timedelta
created_after = utc_now_naive() - timedelta(days=recent_days)

created_after = datetime.now() - timedelta(days=recent_days)
normalized_start_time = normalize_naive_utc_filter_datetime(start_time)
normalized_end_time = normalize_naive_utc_filter_datetime(end_time)

if start_time and end_time and start_time > end_time:
if (
normalized_start_time
and normalized_end_time
and normalized_start_time > normalized_end_time
):
raise ValidationException(
user_message="start_time cannot be later than end_time",
violations=[
{"field": "start_time", "description": "Must be before end_time"}
],
)
# start_time / end_time take priority over recent_days.
if start_time:
created_after = start_time
created_before = end_time
if normalized_start_time:
created_after = normalized_start_time
created_before = normalized_end_time

# Count matching rows.
total_count = await job_repo.count_jobs_by_user(
Expand Down Expand Up @@ -752,10 +767,8 @@ async def list_jobs(

# Compute result_url_expires_at when a download URL was issued.
if result_url:
from datetime import datetime, timedelta

expires_in = int(result_url_info.get("expires_in", 3600))
result_url_expires_at = datetime.now() + timedelta(
result_url_expires_at = utc_now_naive() + timedelta(
seconds=expires_in
)

Expand Down
29 changes: 29 additions & 0 deletions apps/api/tests/contract/test_job_read_contract.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from collections.abc import Callable
from contextlib import AbstractAsyncContextManager
from datetime import datetime, timedelta, timezone
from typing import cast
from uuid import uuid4

Expand Down Expand Up @@ -67,6 +68,34 @@ async def test_should_list_created_jobs_for_the_authenticated_developer(
assert job["credits_spent"] == 0.0


@pytest.mark.asyncio
async def test_should_list_jobs_when_date_filters_use_iso_utc_timezone(
developer_api_client_factory: Callable[
[], AbstractAsyncContextManager[AsyncClient]
],
) -> None:
start_time = datetime.now(timezone.utc) - timedelta(days=1)
end_time = datetime.now(timezone.utc) + timedelta(days=1)
query_params = {
"start_time": start_time.isoformat().replace("+00:00", "Z"),
"end_time": end_time.isoformat().replace("+00:00", "Z"),
}

async with developer_api_client_factory() as api_client:
created_job = await _create_waiting_file_job(api_client)

response = await api_client.get("/api/v1/jobs/page", params=query_params)

assert response.status_code == 200

response_json = cast(dict[str, object], response.json())
jobs = cast(list[dict[str, object]], response_json["jobs"])

assert response_json["total"] == 1
assert len(jobs) == 1
assert jobs[0]["job_id"] == created_job["job_id"]


@pytest.mark.asyncio
async def test_should_return_job_details_for_an_existing_waiting_file_job(
developer_api_client_factory: Callable[
Expand Down
Loading