Skip to content
Open
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
41 changes: 40 additions & 1 deletion api/src/feeds/impl/feeds_api_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,14 +106,26 @@ def get_feeds(
status: str,
provider: str,
producer_url: str,
created_after: str,
created_before: str,
is_official: bool,
db_session: Session,
) -> List[Feed]:
"""Get some (or all) feeds from the Mobility Database."""
is_email_restricted = is_user_email_restricted()
self.logger.debug(f"User email is restricted: {is_email_restricted}")
if created_after and not valid_iso_date(created_after):
raise_http_validation_error(invalid_date_message.format("created_after"))
if created_before and not valid_iso_date(created_before):
raise_http_validation_error(invalid_date_message.format("created_before"))

feed_filter = FeedFilter(
status=status, provider__ilike=provider, producer_url__ilike=producer_url, stable_id=None
status=status,
provider__ilike=provider,
producer_url__ilike=producer_url,
stable_id=None,
created_at__gte=parse_iso_datetime(created_after),
created_at__lte=parse_iso_datetime(created_before),
)
feed_query = feed_filter.filter(Database().get_query_model(db_session, FeedOrm))
feed_query = add_official_filter(feed_query, is_official)
Expand Down Expand Up @@ -199,6 +211,8 @@ def get_gtfs_feeds(
offset: int,
provider: str,
producer_url: str,
created_after: str,
created_before: str,
country_code: str,
subdivision_name: str,
municipality: str,
Expand All @@ -208,13 +222,20 @@ def get_gtfs_feeds(
is_official: bool,
db_session: Session,
) -> List[GtfsFeed]:
if created_after and not valid_iso_date(created_after):
raise_http_validation_error(invalid_date_message.format("created_after"))
if created_before and not valid_iso_date(created_before):
raise_http_validation_error(invalid_date_message.format("created_before"))

try:
published_only = is_user_email_restricted()
feed_query = get_gtfs_feeds_query(
limit=limit,
offset=offset,
provider=provider,
producer_url=producer_url,
created_after=parse_iso_datetime(created_after),
created_before=parse_iso_datetime(created_before),
country_code=country_code,
subdivision_name=subdivision_name,
municipality=municipality,
Expand Down Expand Up @@ -270,6 +291,8 @@ def get_gtfs_rt_feeds(
offset: int,
provider: str,
producer_url: str,
created_after: str,
created_before: str,
entity_types: str,
country_code: str,
subdivision_name: str,
Expand All @@ -278,13 +301,20 @@ def get_gtfs_rt_feeds(
db_session: Session,
) -> List[GtfsRTFeed]:
"""Get some (or all) GTFS Realtime feeds from the Mobility Database."""
if created_after and not valid_iso_date(created_after):
raise_http_validation_error(invalid_date_message.format("created_after"))
if created_before and not valid_iso_date(created_before):
raise_http_validation_error(invalid_date_message.format("created_before"))

try:
published_only = is_user_email_restricted()
feed_query = get_gtfs_rt_feeds_query(
limit=limit,
offset=offset,
provider=provider,
producer_url=producer_url,
created_after=parse_iso_datetime(created_after),
created_before=parse_iso_datetime(created_before),
entity_types=entity_types,
country_code=country_code,
subdivision_name=subdivision_name,
Expand Down Expand Up @@ -545,17 +575,26 @@ def get_gbfs_feeds(
offset: int,
provider: str,
producer_url: str,
created_after: str,
created_before: str,
country_code: str,
subdivision_name: str,
municipality: str,
system_id: str,
version: str,
db_session: Session,
) -> List[GbfsFeed]:
if created_after and not valid_iso_date(created_after):
raise_http_validation_error(invalid_date_message.format("created_after"))
if created_before and not valid_iso_date(created_before):
raise_http_validation_error(invalid_date_message.format("created_before"))

query = get_gbfs_feeds_query(
db_session=db_session,
provider=provider,
producer_url=producer_url,
created_after=parse_iso_datetime(created_after),
created_before=parse_iso_datetime(created_before),
country_code=country_code,
subdivision_name=subdivision_name,
municipality=municipality,
Expand Down
12 changes: 12 additions & 0 deletions api/src/shared/common/db_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,8 @@ def get_gtfs_feeds_query(
offset: int | None = None,
provider: str | None = None,
producer_url: str | None = None,
created_after=None,
created_before=None,
country_code: str | None = None,
subdivision_name: str | None = None,
municipality: str | None = None,
Expand All @@ -56,6 +58,8 @@ def get_gtfs_feeds_query(
stable_id=stable_id,
provider__ilike=provider,
producer_url__ilike=producer_url,
created_at__gte=created_after,
created_at__lte=created_before,
location=None,
)

Expand Down Expand Up @@ -235,6 +239,8 @@ def get_gtfs_rt_feeds_query(
offset: int | None,
provider: str | None,
producer_url: str | None,
created_after,
created_before,
entity_types: str | None,
country_code: str | None,
subdivision_name: str | None,
Expand All @@ -260,6 +266,8 @@ def get_gtfs_rt_feeds_query(
stable_id=None,
provider__ilike=provider,
producer_url__ilike=producer_url,
created_at__gte=created_after,
created_at__lte=created_before,
entity_types=EntityTypeFilter(name__in=entity_types_list),
location=LocationFilter(
country_code=country_code,
Expand Down Expand Up @@ -462,6 +470,8 @@ def get_gbfs_feeds_query(
stable_id: Optional[str] = None,
provider: Optional[str] = None,
producer_url: Optional[str] = None,
created_after=None,
created_before=None,
country_code: Optional[str] = None,
subdivision_name: Optional[str] = None,
municipality: Optional[str] = None,
Expand All @@ -472,6 +482,8 @@ def get_gbfs_feeds_query(
stable_id=stable_id,
provider__ilike=provider,
producer_url__ilike=producer_url,
created_at__gte=created_after,
created_at__lte=created_before,
system_id=system_id,
location=(
LocationFilter(
Expand Down
3 changes: 3 additions & 0 deletions api/src/shared/feed_filters/feed_filter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from typing import Optional
from datetime import datetime

from fastapi_filter.contrib.sqlalchemy import Filter

Expand All @@ -11,6 +12,8 @@ class FeedFilter(Filter):
stable_id: Optional[str]
provider__ilike: Optional[str] # case insensitive
producer_url__ilike: Optional[str] # case insensitive
created_at__gte: Optional[datetime] = None
created_at__lte: Optional[datetime] = None

def __init__(self, *args, **kwargs):
kwargs_normalized = normalize_str_parameter("status", **kwargs)
Expand Down
3 changes: 3 additions & 0 deletions api/src/shared/feed_filters/gbfs_feed_filter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from typing import Optional
from datetime import datetime

from fastapi_filter.contrib.sqlalchemy import Filter

Expand All @@ -17,6 +18,8 @@ class GbfsFeedFilter(Filter):
stable_id: Optional[str] = None
provider__ilike: Optional[str] = None # case-insensitive
producer_url__ilike: Optional[str] = None # case-insensitive
created_at__gte: Optional[datetime] = None
created_at__lte: Optional[datetime] = None
location: Optional[LocationFilter] = None
system_id: Optional[str] = None
version: Optional[GbfsVersionFilter] = None
Expand Down
3 changes: 3 additions & 0 deletions api/src/shared/feed_filters/gtfs_feed_filter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from typing import Optional
from datetime import datetime

from fastapi_filter.contrib.sqlalchemy import Filter

Expand All @@ -25,6 +26,8 @@ class GtfsFeedFilter(Filter):
stable_id: Optional[str]
provider__ilike: Optional[str] # case insensitive
producer_url__ilike: Optional[str] # case insensitive
created_at__gte: Optional[datetime] = None
created_at__lte: Optional[datetime] = None
location: Optional[LocationFilter]

class Constants(Filter.Constants):
Expand Down
3 changes: 3 additions & 0 deletions api/src/shared/feed_filters/gtfs_rt_feed_filter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from typing import Optional, List
from datetime import datetime

from fastapi_filter.contrib.sqlalchemy import Filter

Expand All @@ -22,6 +23,8 @@ class GtfsRtFeedFilter(Filter):
stable_id: Optional[str]
provider__ilike: Optional[str] # case insensitive
producer_url__ilike: Optional[str] # case insensitive
created_at__gte: Optional[datetime] = None
created_at__lte: Optional[datetime] = None
entity_types: Optional[EntityTypeFilter]
location: Optional[LocationFilter]

Expand Down
135 changes: 134 additions & 1 deletion api/tests/integration/test_feeds_api.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
# coding: utf-8
from datetime import datetime, timedelta
import pytest
from fastapi.testclient import TestClient
from datetime import timedelta

from tests.test_utils.database import TEST_GTFS_FEED_STABLE_IDS, TEST_GTFS_RT_FEED_STABLE_ID, TEST_DATASET_STABLE_IDS
from tests.test_utils.token import authHeaders
Expand Down Expand Up @@ -1172,3 +1172,136 @@ def test_gtfs_feed_availability_invalid_date_returns_422(client: TestClient):
params={"from": "not-a-date"},
)
assert response.status_code == 422


@pytest.mark.parametrize(
"endpoint",
[
"/v1/feeds",
"/v1/gtfs_feeds",
"/v1/gtfs_rt_feeds",
"/v1/gbfs_feeds",
],
ids=["feeds", "gtfs_feeds", "gtfs_rt_feeds", "gbfs_feeds"],
)
def test_feed_list_created_at_filters_are_inclusive(client: TestClient, endpoint):
"""Feed-list created_at bounds include feeds exactly on either boundary."""
response = client.request(
"GET",
endpoint,
headers=authHeaders,
)
assert response.status_code == 200

feeds = response.json()
assert feeds, f"Expected test data for {endpoint}"
assert all(feed["created_at"] is not None for feed in feeds)

created_at_by_id = {feed["id"]: datetime.fromisoformat(feed["created_at"].replace("Z", "+00:00")) for feed in feeds}

boundary = max(created_at_by_id.values())
boundary_iso = boundary.isoformat().replace("+00:00", "Z")

expected_ids = {feed_id for feed_id, created_at in created_at_by_id.items() if created_at == boundary}

after_response = client.request(
"GET",
endpoint,
headers=authHeaders,
params={"created_after": boundary_iso},
)
assert after_response.status_code == 200

after_ids = {feed["id"] for feed in after_response.json()}
assert after_ids == expected_ids

before_response = client.request(
"GET",
endpoint,
headers=authHeaders,
params={"created_before": boundary_iso},
)
assert before_response.status_code == 200

before_ids = {feed["id"] for feed in before_response.json()}
assert expected_ids.issubset(before_ids)


@pytest.mark.parametrize(
"endpoint",
[
"/v1/feeds",
"/v1/gtfs_feeds",
"/v1/gtfs_rt_feeds",
"/v1/gbfs_feeds",
],
ids=["feeds", "gtfs_feeds", "gtfs_rt_feeds", "gbfs_feeds"],
)
def test_feed_list_created_at_range(client: TestClient, endpoint):
"""Combined created_at bounds return exactly the inclusive temporal range."""
response = client.request(
"GET",
endpoint,
headers=authHeaders,
)
assert response.status_code == 200

feeds = response.json()
assert feeds, f"Expected test data for {endpoint}"
assert all(feed["created_at"] is not None for feed in feeds)

created_at_by_id = {feed["id"]: datetime.fromisoformat(feed["created_at"].replace("Z", "+00:00")) for feed in feeds}

ordered_dates = sorted(set(created_at_by_id.values()))

lower = ordered_dates[0]
upper = ordered_dates[-1]

lower_iso = lower.isoformat().replace("+00:00", "Z")
upper_iso = upper.isoformat().replace("+00:00", "Z")

expected_ids = {feed_id for feed_id, created_at in created_at_by_id.items() if lower <= created_at <= upper}

filtered_response = client.request(
"GET",
endpoint,
headers=authHeaders,
params={
"created_after": lower_iso,
"created_before": upper_iso,
},
)
assert filtered_response.status_code == 200

filtered_ids = {feed["id"] for feed in filtered_response.json()}
assert filtered_ids == expected_ids


@pytest.mark.parametrize(
"endpoint",
[
"/v1/feeds",
"/v1/gtfs_feeds",
"/v1/gtfs_rt_feeds",
"/v1/gbfs_feeds",
],
ids=["feeds", "gtfs_feeds", "gtfs_rt_feeds", "gbfs_feeds"],
)
@pytest.mark.parametrize(
"parameter",
["created_after", "created_before"],
)
def test_feed_list_created_at_invalid_date_returns_422(
client: TestClient,
endpoint,
parameter,
):
"""Invalid created_at bounds are rejected consistently by feed-list endpoints."""
response = client.request(
"GET",
endpoint,
headers=authHeaders,
params={parameter: "invalid_date"},
)

assert response.status_code == 422
Loading
Loading