From 49fa4aeef7ac7e2c5901c0d4fd58e127afa09342 Mon Sep 17 00:00:00 2001 From: Vincent WENDLING Date: Sat, 18 Apr 2026 22:29:28 +0200 Subject: [PATCH 1/3] Add unit tests --- src/trendflow/_trends_http/transport.py | 2 +- tests/conftest.py | 126 ++++++ tests/test_enums.py | 90 +++++ tests/test_exceptions.py | 79 ++++ tests/test_exporters.py | 192 ++++++++++ tests/test_fetcher.py | 281 ++++++++++++++ tests/test_models.py | 212 ++++++++++ tests/test_parsers.py | 489 ++++++++++++++++++++++++ tests/test_session.py | 223 +++++++++++ tests/test_transport.py | 270 +++++++++++++ tests/test_trendflow.py | 49 ++- 11 files changed, 2011 insertions(+), 2 deletions(-) create mode 100644 tests/conftest.py create mode 100644 tests/test_enums.py create mode 100644 tests/test_exceptions.py create mode 100644 tests/test_exporters.py create mode 100644 tests/test_fetcher.py create mode 100644 tests/test_models.py create mode 100644 tests/test_parsers.py create mode 100644 tests/test_session.py create mode 100644 tests/test_transport.py diff --git a/src/trendflow/_trends_http/transport.py b/src/trendflow/_trends_http/transport.py index 2b8b248..fede18e 100644 --- a/src/trendflow/_trends_http/transport.py +++ b/src/trendflow/_trends_http/transport.py @@ -21,7 +21,7 @@ def _normalize_timeout(timeout: httpx.Timeout | tuple[float, float] | float) -> return timeout if isinstance(timeout, tuple): a, b = float(timeout[0]), float(timeout[1]) - return httpx.Timeout(connect=a, read=b) + return httpx.Timeout(b, connect=a, read=b) return httpx.Timeout(float(timeout)) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..527dc68 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,126 @@ +"""Shared fixtures for the trendflow test suite.""" + +from __future__ import annotations + +from datetime import datetime + +import pytest + +from trendflow.enums import ExportFormat, Region, Resolution, Timeframe +from trendflow.models import ( + InterestByRegionResult, + InterestOverTimeResult, + RegionalInterestRow, + RelatedQuery, + RelatedResult, + TrendingItem, + TrendingResult, + TrendPoint, +) + + +@pytest.fixture +def dt_jan1() -> datetime: + return datetime(2024, 1, 1, 0, 0, 0) + + +@pytest.fixture +def dt_jan8() -> datetime: + return datetime(2024, 1, 8, 0, 0, 0) + + +@pytest.fixture +def two_kw_points() -> list[TrendPoint]: + return [ + TrendPoint(date=datetime(2024, 1, 1), scores={"Python": 80, "JavaScript": 70}), + TrendPoint(date=datetime(2024, 1, 8), scores={"Python": 85, "JavaScript": 65}), + TrendPoint(date=datetime(2024, 1, 15), scores={"Python": 90, "JavaScript": 60}), + ] + + +@pytest.fixture +def iot_result(two_kw_points: list[TrendPoint]) -> InterestOverTimeResult: + return InterestOverTimeResult( + keywords=["Python", "JavaScript"], + granularity="weekly", + points=two_kw_points, + ) + + +@pytest.fixture +def empty_iot_result() -> InterestOverTimeResult: + return InterestOverTimeResult(keywords=["Python"], granularity="unknown", points=[]) + + +@pytest.fixture +def region_rows() -> list[RegionalInterestRow]: + return [ + RegionalInterestRow(label="California", value=90), + RegionalInterestRow(label="Texas", value=70), + ] + + +@pytest.fixture +def ibr_result(region_rows: list[RegionalInterestRow]) -> InterestByRegionResult: + return InterestByRegionResult( + keyword="Python", resolution=Resolution.REGION, rows=region_rows + ) + + +@pytest.fixture +def trending_result() -> TrendingResult: + return TrendingResult( + results=[ + TrendingItem(title="AI tools", traffic="500K+", articles=[]), + TrendingItem(title="Python 4", traffic="200K+", articles=[]), + ] + ) + + +@pytest.fixture +def related_result() -> RelatedResult: + return RelatedResult( + top=[RelatedQuery(term="python tutorial", value=100)], + rising=[RelatedQuery(term="python ai", breakout="+250%")], + ) + + +@pytest.fixture +def timeline_data_weekly() -> dict: + """Two entries 7 days apart for granularity inference.""" + ts0 = int(datetime(2024, 1, 1).timestamp()) + ts1 = int(datetime(2024, 1, 8).timestamp()) + ts2 = int(datetime(2024, 1, 15).timestamp()) + return { + "timelineData": [ + {"time": str(ts0), "value": "[80, 70]"}, + {"time": str(ts1), "value": "[85, 65]"}, + {"time": str(ts2), "value": "[90, 60]"}, + ] + } + + +@pytest.fixture +def timeline_data_daily() -> dict: + """Two entries 1 day apart.""" + ts0 = int(datetime(2024, 1, 1).timestamp()) + ts1 = int(datetime(2024, 1, 2).timestamp()) + return { + "timelineData": [ + {"time": str(ts0), "value": "[50]"}, + {"time": str(ts1), "value": "[55]"}, + ] + } + + +@pytest.fixture +def timeline_data_hourly() -> dict: + """Two entries 1 hour apart.""" + ts0 = int(datetime(2024, 1, 1, 0, 0, 0).timestamp()) + ts1 = int(datetime(2024, 1, 1, 1, 0, 0).timestamp()) + return { + "timelineData": [ + {"time": str(ts0), "value": "[40]"}, + {"time": str(ts1), "value": "[42]"}, + ] + } diff --git a/tests/test_enums.py b/tests/test_enums.py new file mode 100644 index 0000000..5209d83 --- /dev/null +++ b/tests/test_enums.py @@ -0,0 +1,90 @@ +"""Tests for trendflow.enums.""" + +from __future__ import annotations + +import pytest + +from trendflow.enums import ExportFormat, Region, Resolution, Timeframe + + +class TestRegion: + def test_worldwide_is_empty_string(self) -> None: + assert Region.WORLDWIDE == "" + + def test_us_value(self) -> None: + assert Region.US == "US" + + def test_all_non_worldwide_are_uppercase_two_letter(self) -> None: + for region in Region: + if region is not Region.WORLDWIDE: + assert len(region.value) == 2 + assert region.value.isupper() + + def test_str_serialization(self) -> None: + assert str(Region.GB) == "GB" + assert str(Region.DE) == "DE" + + def test_all_expected_regions_present(self) -> None: + codes = {r.value for r in Region} + for expected in ("US", "GB", "DE", "FR", "IT", "ES", "CA", "AU", "JP", "IN", "BR", "MX"): + assert expected in codes + + def test_comparable_to_string(self) -> None: + assert Region.US == "US" + assert "US" == Region.US + + def test_usable_in_f_string(self) -> None: + assert f"geo={Region.US}" == "geo=US" + assert f"geo={Region.WORLDWIDE}" == "geo=" + + +class TestTimeframe: + def test_past_day_value(self) -> None: + assert Timeframe.PAST_DAY == "now 1-d" + + def test_past_week_value(self) -> None: + assert Timeframe.PAST_WEEK == "now 7-d" + + def test_past_year_value(self) -> None: + assert Timeframe.PAST_YEAR == "today 12-m" + + def test_past_5_years_value(self) -> None: + assert Timeframe.PAST_5_YEARS == "today 5-y" + + def test_all_four_timeframes_exist(self) -> None: + assert len(list(Timeframe)) == 4 + + def test_str_serialization(self) -> None: + assert str(Timeframe.PAST_DAY) == "now 1-d" + + +class TestResolution: + def test_country_value(self) -> None: + assert Resolution.COUNTRY == "COUNTRY" + + def test_region_value(self) -> None: + assert Resolution.REGION == "REGION" + + def test_city_value(self) -> None: + assert Resolution.CITY == "CITY" + + def test_all_three_exist(self) -> None: + assert len(list(Resolution)) == 3 + + def test_str_serialization(self) -> None: + assert str(Resolution.CITY) == "CITY" + + +class TestExportFormat: + def test_csv_value(self) -> None: + assert ExportFormat.CSV == "csv" + + def test_json_value(self) -> None: + assert ExportFormat.JSON == "json" + + def test_both_formats_exist(self) -> None: + assert len(list(ExportFormat)) == 2 + + def test_str_serialization(self) -> None: + assert str(ExportFormat.CSV) == "csv" + assert str(ExportFormat.JSON) == "json" diff --git a/tests/test_exceptions.py b/tests/test_exceptions.py new file mode 100644 index 0000000..37a22b4 --- /dev/null +++ b/tests/test_exceptions.py @@ -0,0 +1,79 @@ +"""Tests for trendflow._trends_http.exceptions.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import httpx +import pytest + +from trendflow._trends_http.exceptions import ResponseError, TooManyRequestsError + + +def _make_response(status_code: int) -> httpx.Response: + mock = MagicMock(spec=httpx.Response) + mock.status_code = status_code + return mock + + +class TestResponseError: + def test_message_in_args(self) -> None: + response = _make_response(500) + err = ResponseError("something went wrong", response) + assert "something went wrong" in str(err) + + def test_response_attached(self) -> None: + response = _make_response(500) + err = ResponseError("fail", response) + assert err.response is response + + def test_is_exception(self) -> None: + response = _make_response(500) + err = ResponseError("fail", response) + assert isinstance(err, Exception) + + def test_from_response_classmethod(self) -> None: + response = _make_response(503) + err = ResponseError.from_response(response) + assert isinstance(err, ResponseError) + assert "503" in str(err) + assert err.response is response + + def test_from_response_message_format(self) -> None: + response = _make_response(404) + err = ResponseError.from_response(response) + assert "404" in str(err) + assert "Google" in str(err) + + def test_can_be_raised_and_caught(self) -> None: + response = _make_response(500) + with pytest.raises(ResponseError) as exc_info: + raise ResponseError("test error", response) + assert exc_info.value.response is response + + +class TestTooManyRequestsError: + def test_is_response_error_subclass(self) -> None: + response = _make_response(429) + err = TooManyRequestsError("rate limited", response) + assert isinstance(err, ResponseError) + + def test_is_exception(self) -> None: + response = _make_response(429) + err = TooManyRequestsError("rate limited", response) + assert isinstance(err, Exception) + + def test_from_response_returns_too_many_requests_error(self) -> None: + response = _make_response(429) + err = TooManyRequestsError.from_response(response) + assert isinstance(err, TooManyRequestsError) + + def test_response_attached(self) -> None: + response = _make_response(429) + err = TooManyRequestsError("rate limited", response) + assert err.response is response + + def test_can_catch_as_response_error(self) -> None: + response = _make_response(429) + with pytest.raises(ResponseError): + raise TooManyRequestsError("rate limited", response) diff --git a/tests/test_exporters.py b/tests/test_exporters.py new file mode 100644 index 0000000..459a618 --- /dev/null +++ b/tests/test_exporters.py @@ -0,0 +1,192 @@ +"""Tests for trendflow._exporters.""" + +from __future__ import annotations + +import csv +import json +import tempfile +from datetime import datetime +from pathlib import Path + +import pytest + +from trendflow._exporters import ( + INTEREST_OVER_TIME_EXPORTERS, + CsvInterestOverTimeExporter, + JsonInterestOverTimeExporter, + export_interest_over_time, +) +from trendflow.enums import ExportFormat +from trendflow.models import InterestOverTimeResult, TrendPoint + + +@pytest.fixture +def result_with_points() -> InterestOverTimeResult: + return InterestOverTimeResult( + keywords=["Python", "JavaScript"], + granularity="weekly", + points=[ + TrendPoint(date=datetime(2024, 1, 1), scores={"Python": 80, "JavaScript": 70}), + TrendPoint(date=datetime(2024, 1, 8), scores={"Python": 85, "JavaScript": 65}), + ], + ) + + +@pytest.fixture +def result_empty() -> InterestOverTimeResult: + return InterestOverTimeResult(keywords=["Python"], granularity="unknown", points=[]) + + +class TestCsvExporter: + def test_creates_file(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_with_points, path) + assert path.exists() + finally: + path.unlink(missing_ok=True) + + def test_csv_has_header_row(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_with_points, path) + text = path.read_text(encoding="utf-8") + first_line = text.splitlines()[0] + assert "date" in first_line + assert "Python" in first_line + assert "JavaScript" in first_line + finally: + path.unlink(missing_ok=True) + + def test_csv_row_count_matches_points(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_with_points, path) + with path.open(encoding="utf-8") as fh: + rows = list(csv.DictReader(fh)) + assert len(rows) == 2 + finally: + path.unlink(missing_ok=True) + + def test_csv_values_present(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_with_points, path) + text = path.read_text(encoding="utf-8") + assert "80" in text + assert "70" in text + finally: + path.unlink(missing_ok=True) + + def test_csv_empty_result_has_only_header(self, result_empty: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_empty, path) + with path.open(encoding="utf-8") as fh: + rows = list(csv.DictReader(fh)) + assert rows == [] + finally: + path.unlink(missing_ok=True) + + def test_csv_utf8_encoding(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + CsvInterestOverTimeExporter().export(result_with_points, path) + # Should not raise + path.read_text(encoding="utf-8") + finally: + path.unlink(missing_ok=True) + + +class TestJsonExporter: + def test_creates_file(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + JsonInterestOverTimeExporter().export(result_with_points, path) + assert path.exists() + finally: + path.unlink(missing_ok=True) + + def test_json_is_valid(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + JsonInterestOverTimeExporter().export(result_with_points, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert isinstance(data, list) + finally: + path.unlink(missing_ok=True) + + def test_json_record_count_matches_points(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + JsonInterestOverTimeExporter().export(result_with_points, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert len(data) == 2 + finally: + path.unlink(missing_ok=True) + + def test_json_records_have_keyword_keys(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + JsonInterestOverTimeExporter().export(result_with_points, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert "Python" in data[0] + assert "JavaScript" in data[0] + finally: + path.unlink(missing_ok=True) + + def test_json_empty_result_is_empty_list(self, result_empty: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + JsonInterestOverTimeExporter().export(result_empty, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert data == [] + finally: + path.unlink(missing_ok=True) + + +class TestExportInterestOverTime: + def test_dispatches_csv(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + export_interest_over_time(result_with_points, ExportFormat.CSV, path) + text = path.read_text(encoding="utf-8") + assert "Python" in text + finally: + path.unlink(missing_ok=True) + + def test_dispatches_json(self, result_with_points: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + export_interest_over_time(result_with_points, ExportFormat.JSON, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert isinstance(data, list) + finally: + path.unlink(missing_ok=True) + + def test_unsupported_format_raises_value_error(self, result_with_points: InterestOverTimeResult) -> None: + with pytest.raises(ValueError, match="Unsupported export format"): + export_interest_over_time(result_with_points, "xml", Path("/tmp/out.xml")) # type: ignore[arg-type] + + def test_exporters_registry_has_csv_and_json(self) -> None: + assert ExportFormat.CSV in INTEREST_OVER_TIME_EXPORTERS + assert ExportFormat.JSON in INTEREST_OVER_TIME_EXPORTERS + + def test_csv_exporter_type(self) -> None: + assert isinstance(INTEREST_OVER_TIME_EXPORTERS[ExportFormat.CSV], CsvInterestOverTimeExporter) + + def test_json_exporter_type(self) -> None: + assert isinstance(INTEREST_OVER_TIME_EXPORTERS[ExportFormat.JSON], JsonInterestOverTimeExporter) diff --git a/tests/test_fetcher.py b/tests/test_fetcher.py new file mode 100644 index 0000000..4611b66 --- /dev/null +++ b/tests/test_fetcher.py @@ -0,0 +1,281 @@ +"""Tests for trendflow._fetcher (GoogleTrendsFetcher and helpers).""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from trendflow._fetcher import ( + TRENDING_PN, + GoogleTrendsFetcher, + TrendsFetcher, + _hl_from_language, +) +from trendflow.enums import Region, Resolution, Timeframe +from trendflow.models import ( + InterestByRegionResult, + InterestOverTimeResult, + RelatedResult, + TrendingResult, +) + + +class TestHlFromLanguage: + def test_language_without_dash_appends_us(self) -> None: + assert _hl_from_language("en") == "en-US" + + def test_language_with_dash_returned_as_is(self) -> None: + assert _hl_from_language("en-GB") == "en-GB" + + def test_two_char_language_gets_us_suffix(self) -> None: + assert _hl_from_language("fr") == "fr-US" + + def test_already_full_locale_unchanged(self) -> None: + assert _hl_from_language("zh-CN") == "zh-CN" + + +class TestTrendingPnMapping: + def test_us_maps_to_united_states(self) -> None: + assert TRENDING_PN[Region.US] == "united_states" + + def test_gb_maps_to_united_kingdom(self) -> None: + assert TRENDING_PN[Region.GB] == "united_kingdom" + + def test_worldwide_not_in_mapping(self) -> None: + assert Region.WORLDWIDE not in TRENDING_PN + + def test_all_non_worldwide_regions_have_mapping(self) -> None: + for region in Region: + if region is not Region.WORLDWIDE: + assert region in TRENDING_PN, f"Missing TRENDING_PN entry for {region!r}" + + +def _make_fetcher() -> GoogleTrendsFetcher: + """Build a GoogleTrendsFetcher with GoogleTrendsHttpSession patched out.""" + with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_session_cls: + mock_session = MagicMock() + mock_session.geo = "US" + mock_session_cls.return_value = mock_session + fetcher = GoogleTrendsFetcher() + return fetcher + + +class TestGoogleTrendsFetcherInit: + def test_default_language_passed_as_hl(self) -> None: + with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_cls: + mock_cls.return_value = MagicMock(geo="") + GoogleTrendsFetcher(language="en") + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["hl"] == "en-US" + + def test_hyphenated_language_passed_unchanged(self) -> None: + with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_cls: + mock_cls.return_value = MagicMock(geo="") + GoogleTrendsFetcher(language="en-GB") + call_kwargs = mock_cls.call_args.kwargs + assert call_kwargs["hl"] == "en-GB" + + def test_timeout_tuple_uses_doubled_read(self) -> None: + with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_cls: + mock_cls.return_value = MagicMock(geo="") + GoogleTrendsFetcher(timeout=10) + call_kwargs = mock_cls.call_args.kwargs + connect, read = call_kwargs["timeout"] + assert connect == 10 + assert read == max(20, 15) # max(timeout*2, timeout+5) = max(20, 15) = 20 + + def test_timeout_min_read_is_timeout_plus_5(self) -> None: + with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_cls: + mock_cls.return_value = MagicMock(geo="") + GoogleTrendsFetcher(timeout=3) + call_kwargs = mock_cls.call_args.kwargs + connect, read = call_kwargs["timeout"] + assert connect == 3 + assert read == max(6, 8) # max(3*2, 3+5) = max(6, 8) = 8 + + +class TestTrendsFetcherProtocol: + def test_google_trends_fetcher_implements_protocol(self) -> None: + fetcher = _make_fetcher() + assert isinstance(fetcher, TrendsFetcher) + + def test_protocol_methods_exist(self) -> None: + fetcher = _make_fetcher() + assert hasattr(fetcher, "interest_over_time") + assert hasattr(fetcher, "interest_by_region") + assert hasattr(fetcher, "trending_now") + assert hasattr(fetcher, "related_queries") + + +class TestInterestOverTime: + def test_calls_build_payload_with_correct_args(self) -> None: + fetcher = _make_fetcher() + mock_default = {"timelineData": []} + fetcher._req.interest_over_time.return_value = mock_default + fetcher._req.geo = "US" + + fetcher.interest_over_time( + keywords=["Python"], + timeframe=Timeframe.PAST_YEAR, + region=Region.US, + ) + + fetcher._req.build_payload.assert_called_once_with( + ["Python"], + cat=0, + timeframe=Timeframe.PAST_YEAR.value, + geo=Region.US.value, + gprop="", + ) + + def test_returns_interest_over_time_result(self) -> None: + fetcher = _make_fetcher() + fetcher._req.interest_over_time.return_value = {"timelineData": []} + fetcher._req.geo = "US" + + result = fetcher.interest_over_time( + keywords=["Python"], + timeframe=Timeframe.PAST_YEAR, + region=Region.US, + ) + + assert isinstance(result, InterestOverTimeResult) + + def test_passes_geo_from_session(self) -> None: + fetcher = _make_fetcher() + fetcher._req.interest_over_time.return_value = {"timelineData": []} + fetcher._req.geo = ["US"] + + result = fetcher.interest_over_time( + keywords=["Python"], + timeframe=Timeframe.PAST_YEAR, + region=Region.US, + ) + + assert result.keywords == ["Python"] + + +class TestInterestByRegion: + def test_returns_empty_result_when_no_geo_map_data(self) -> None: + fetcher = _make_fetcher() + fetcher._req.interest_by_region.return_value = {} + + result = fetcher.interest_by_region( + keyword="Python", + resolution=Resolution.COUNTRY, + region=Region.US, + ) + + assert isinstance(result, InterestByRegionResult) + assert result.rows == [] + assert result.keyword == "Python" + + def test_calls_build_payload_with_keyword(self) -> None: + fetcher = _make_fetcher() + fetcher._req.interest_by_region.return_value = {} + + fetcher.interest_by_region(keyword="Rust", resolution=Resolution.REGION, region=Region.US) + + fetcher._req.build_payload.assert_called_once_with( + ["Rust"], + cat=0, + timeframe=Timeframe.PAST_YEAR.value, + geo=Region.US.value, + gprop="", + ) + + def test_returns_parsed_result_when_data_present(self) -> None: + fetcher = _make_fetcher() + fetcher._req.interest_by_region.return_value = { + "geoMapData": [{"geoName": "California", "value": "[90]"}] + } + + result = fetcher.interest_by_region( + keyword="Python", resolution=Resolution.REGION, region=Region.US + ) + + assert isinstance(result, InterestByRegionResult) + assert len(result.rows) == 1 + assert result.rows[0].label == "California" + + +class TestTrendingNow: + def test_worldwide_raises_value_error(self) -> None: + fetcher = _make_fetcher() + with pytest.raises(ValueError, match="specific country"): + fetcher.trending_now(region=Region.WORLDWIDE) + + def test_valid_region_calls_trending_searches(self) -> None: + fetcher = _make_fetcher() + fetcher._req.trending_searches.return_value = ["AI", "Python"] + + result = fetcher.trending_now(region=Region.US) + + fetcher._req.trending_searches.assert_called_once_with(pn="united_states") + assert isinstance(result, TrendingResult) + + def test_returns_correct_titles(self) -> None: + fetcher = _make_fetcher() + fetcher._req.trending_searches.return_value = ["AI tools", "Python 4"] + + result = fetcher.trending_now(region=Region.US) + + assert result.results[0].title == "AI tools" + assert result.results[1].title == "Python 4" + + def test_pn_lookup_for_gb(self) -> None: + fetcher = _make_fetcher() + fetcher._req.trending_searches.return_value = [] + + fetcher.trending_now(region=Region.GB) + + fetcher._req.trending_searches.assert_called_once_with(pn="united_kingdom") + + def test_pn_lookup_for_de(self) -> None: + fetcher = _make_fetcher() + fetcher._req.trending_searches.return_value = [] + + fetcher.trending_now(region=Region.DE) + + fetcher._req.trending_searches.assert_called_once_with(pn="germany") + + +class TestRelatedQueries: + def test_calls_build_payload_with_keyword(self) -> None: + fetcher = _make_fetcher() + fetcher._req.related_queries.return_value = {} + + fetcher.related_queries(keyword="Python") + + fetcher._req.build_payload.assert_called_once_with( + ["Python"], + cat=0, + timeframe=Timeframe.PAST_YEAR.value, + geo="", + gprop="", + ) + + def test_returns_related_result(self) -> None: + fetcher = _make_fetcher() + fetcher._req.related_queries.return_value = { + "Python": { + "top": [{"query": "python tutorial", "value": 100}], + "rising": [], + } + } + + result = fetcher.related_queries(keyword="Python") + + assert isinstance(result, RelatedResult) + assert len(result.top) == 1 + assert result.top[0].term == "python tutorial" + + def test_empty_raw_returns_empty_result(self) -> None: + fetcher = _make_fetcher() + fetcher._req.related_queries.return_value = {} + + result = fetcher.related_queries(keyword="Python") + + assert result.top == [] + assert result.rising == [] diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 0000000..f2f321d --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,212 @@ +"""Tests for trendflow.models.""" + +from __future__ import annotations + +import json +import tempfile +from datetime import datetime +from pathlib import Path + +import pandas as pd +import pytest + +from trendflow.enums import ExportFormat, Resolution +from trendflow.models import ( + InterestByRegionResult, + InterestOverTimeResult, + RegionalInterestRow, + RelatedQuery, + RelatedResult, + TrendingItem, + TrendingResult, + TrendPoint, +) + + +class TestTrendPoint: + def test_construction(self) -> None: + dt = datetime(2024, 1, 1) + point = TrendPoint(date=dt, scores={"Python": 80}) + assert point.date == dt + assert point.scores == {"Python": 80} + + def test_frozen(self) -> None: + point = TrendPoint(date=datetime(2024, 1, 1), scores={"Python": 80}) + with pytest.raises(Exception): + point.date = datetime(2024, 1, 2) # type: ignore[misc] + + def test_equality(self) -> None: + dt = datetime(2024, 1, 1) + p1 = TrendPoint(date=dt, scores={"Python": 80}) + p2 = TrendPoint(date=dt, scores={"Python": 80}) + assert p1 == p2 + + def test_multiple_keywords_in_scores(self) -> None: + point = TrendPoint( + date=datetime(2024, 1, 1), + scores={"Python": 80, "JavaScript": 70, "Rust": 50}, + ) + assert point.scores["Python"] == 80 + assert point.scores["Rust"] == 50 + + +class TestInterestOverTimeResult: + def test_to_dataframe_empty(self, empty_iot_result: InterestOverTimeResult) -> None: + df = empty_iot_result.to_dataframe() + assert isinstance(df, pd.DataFrame) + assert len(df) == 0 + assert "date" in df.columns + assert "Python" in df.columns + + def test_to_dataframe_empty_has_keyword_columns(self) -> None: + result = InterestOverTimeResult( + keywords=["A", "B", "C"], granularity="unknown", points=[] + ) + df = result.to_dataframe() + assert list(df.columns) == ["date", "A", "B", "C"] + + def test_to_dataframe_with_points(self, iot_result: InterestOverTimeResult) -> None: + df = iot_result.to_dataframe() + assert isinstance(df, pd.DataFrame) + assert len(df) == 3 + assert "date" in df.columns + assert "Python" in df.columns + assert "JavaScript" in df.columns + + def test_to_dataframe_values(self, iot_result: InterestOverTimeResult) -> None: + df = iot_result.to_dataframe() + assert df["Python"].tolist() == [80, 85, 90] + assert df["JavaScript"].tolist() == [70, 65, 60] + + def test_to_dataframe_date_column(self, iot_result: InterestOverTimeResult) -> None: + df = iot_result.to_dataframe() + assert df["date"].iloc[0] == datetime(2024, 1, 1) + + def test_frozen(self, iot_result: InterestOverTimeResult) -> None: + with pytest.raises(Exception): + iot_result.keywords = ["other"] # type: ignore[misc] + + def test_export_csv(self, iot_result: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = Path(f.name) + try: + iot_result.export(ExportFormat.CSV, path) + content = path.read_text(encoding="utf-8") + assert "Python" in content + assert "JavaScript" in content + assert "80" in content + finally: + path.unlink(missing_ok=True) + + def test_export_json(self, iot_result: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: + path = Path(f.name) + try: + iot_result.export(ExportFormat.JSON, path) + data = json.loads(path.read_text(encoding="utf-8")) + assert isinstance(data, list) + assert len(data) == 3 + finally: + path.unlink(missing_ok=True) + + def test_export_accepts_string_path(self, iot_result: InterestOverTimeResult) -> None: + with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: + path = f.name + try: + iot_result.export(ExportFormat.CSV, path) + assert Path(path).exists() + finally: + Path(path).unlink(missing_ok=True) + + def test_single_keyword_dataframe(self) -> None: + result = InterestOverTimeResult( + keywords=["Rust"], + granularity="daily", + points=[TrendPoint(date=datetime(2024, 1, 1), scores={"Rust": 42})], + ) + df = result.to_dataframe() + assert list(df.columns) == ["date", "Rust"] + assert df["Rust"].iloc[0] == 42 + + +class TestRegionalInterestRow: + def test_construction(self) -> None: + row = RegionalInterestRow(label="California", value=90) + assert row.label == "California" + assert row.value == 90 + + def test_frozen(self) -> None: + row = RegionalInterestRow(label="California", value=90) + with pytest.raises(Exception): + row.value = 100 # type: ignore[misc] + + +class TestInterestByRegionResult: + def test_construction(self, ibr_result: InterestByRegionResult) -> None: + assert ibr_result.keyword == "Python" + assert ibr_result.resolution == Resolution.REGION + assert len(ibr_result.rows) == 2 + + def test_frozen(self, ibr_result: InterestByRegionResult) -> None: + with pytest.raises(Exception): + ibr_result.keyword = "other" # type: ignore[misc] + + def test_empty_rows(self) -> None: + result = InterestByRegionResult( + keyword="Python", resolution=Resolution.COUNTRY, rows=[] + ) + assert result.rows == [] + + +class TestTrendingItem: + def test_construction(self) -> None: + item = TrendingItem(title="AI news", traffic="500K+", articles=[]) + assert item.title == "AI news" + assert item.traffic == "500K+" + assert item.articles == [] + + def test_frozen(self) -> None: + item = TrendingItem(title="AI news", traffic="500K+", articles=[]) + with pytest.raises(Exception): + item.title = "other" # type: ignore[misc] + + +class TestTrendingResult: + def test_construction(self, trending_result: TrendingResult) -> None: + assert len(trending_result.results) == 2 + assert trending_result.results[0].title == "AI tools" + + def test_empty(self) -> None: + result = TrendingResult(results=[]) + assert result.results == [] + + +class TestRelatedQuery: + def test_defaults(self) -> None: + q = RelatedQuery(term="python tutorial") + assert q.value is None + assert q.breakout is None + + def test_with_value(self) -> None: + q = RelatedQuery(term="python tutorial", value=100) + assert q.value == 100 + + def test_with_breakout(self) -> None: + q = RelatedQuery(term="python ai", breakout="+250%") + assert q.breakout == "+250%" + + def test_frozen(self) -> None: + q = RelatedQuery(term="python tutorial", value=100) + with pytest.raises(Exception): + q.term = "other" # type: ignore[misc] + + +class TestRelatedResult: + def test_construction(self, related_result: RelatedResult) -> None: + assert len(related_result.top) == 1 + assert len(related_result.rising) == 1 + + def test_empty(self) -> None: + result = RelatedResult(top=[], rising=[]) + assert result.top == [] + assert result.rising == [] diff --git a/tests/test_parsers.py b/tests/test_parsers.py new file mode 100644 index 0000000..8b92c89 --- /dev/null +++ b/tests/test_parsers.py @@ -0,0 +1,489 @@ +"""Tests for trendflow._parsers.""" + +from __future__ import annotations + +import math +from datetime import datetime, timedelta + +import pytest + +from trendflow._parsers import ( + _is_missing_value, + _split_bracketed_ints, + _to_int_or_none, + infer_granularity, + interest_by_region_rows, + interest_by_region_to_result, + interest_over_time_to_result, + parse_rising_related, + parse_top_related, + related_queries_to_result, + trending_result_from_titles, + trending_titles_to_items, +) +from trendflow.enums import Resolution +from trendflow.models import ( + InterestByRegionResult, + InterestOverTimeResult, + RegionalInterestRow, + RelatedQuery, + RelatedResult, + TrendingItem, + TrendingResult, +) + + +class TestSplitBracketedInts: + def test_bracketed_pair(self) -> None: + assert _split_bracketed_ints("[80, 70]") == [80, 70] + + def test_single_value(self) -> None: + assert _split_bracketed_ints("[42]") == [42] + + def test_no_brackets(self) -> None: + assert _split_bracketed_ints("55") == [55] + + def test_extra_spaces(self) -> None: + assert _split_bracketed_ints("[ 10 , 20 ]") == [10, 20] + + def test_empty_string(self) -> None: + assert _split_bracketed_ints("") == [] + + def test_empty_brackets(self) -> None: + assert _split_bracketed_ints("[]") == [] + + def test_three_values(self) -> None: + assert _split_bracketed_ints("[1, 2, 3]") == [1, 2, 3] + + def test_integer_input(self) -> None: + assert _split_bracketed_ints(99) == [99] + + +class TestIsMissingValue: + def test_none_is_missing(self) -> None: + assert _is_missing_value(None) is True + + def test_nan_is_missing(self) -> None: + assert _is_missing_value(float("nan")) is True + + def test_math_nan_is_missing(self) -> None: + assert _is_missing_value(math.nan) is True + + def test_zero_is_not_missing(self) -> None: + assert _is_missing_value(0) is False + + def test_integer_not_missing(self) -> None: + assert _is_missing_value(42) is False + + def test_string_not_missing(self) -> None: + assert _is_missing_value("hello") is False + + def test_empty_string_not_missing(self) -> None: + assert _is_missing_value("") is False + + def test_false_not_missing(self) -> None: + assert _is_missing_value(False) is False + + def test_float_not_nan_not_missing(self) -> None: + assert _is_missing_value(3.14) is False + + +class TestInferGranularity: + def test_7_days_is_weekly(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = datetime(2024, 1, 8) + assert infer_granularity(d0, d1) == "weekly" + + def test_6_days_is_weekly(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = datetime(2024, 1, 7) + assert infer_granularity(d0, d1) == "weekly" + + def test_1_day_is_daily(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = datetime(2024, 1, 2) + assert infer_granularity(d0, d1) == "daily" + + def test_5_days_is_daily(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = datetime(2024, 1, 6) + assert infer_granularity(d0, d1) == "daily" + + def test_1_hour_is_hourly(self) -> None: + d0 = datetime(2024, 1, 1, 0, 0) + d1 = datetime(2024, 1, 1, 1, 0) + assert infer_granularity(d0, d1) == "hourly" + + def test_30_minutes_is_hourly(self) -> None: + d0 = datetime(2024, 1, 1, 0, 0) + d1 = datetime(2024, 1, 1, 0, 30) + assert infer_granularity(d0, d1) == "hourly" + + def test_boundary_exactly_6_days(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = d0 + timedelta(days=6) + assert infer_granularity(d0, d1) == "weekly" + + def test_boundary_exactly_1_day(self) -> None: + d0 = datetime(2024, 1, 1) + d1 = d0 + timedelta(days=1) + assert infer_granularity(d0, d1) == "daily" + + +class TestInterestOverTimeToResult: + def test_empty_timeline(self) -> None: + result = interest_over_time_to_result({}, ["Python"], "US") + assert isinstance(result, InterestOverTimeResult) + assert result.granularity == "unknown" + assert result.points == [] + assert result.keywords == ["Python"] + + def test_none_timeline(self) -> None: + result = interest_over_time_to_result({"timelineData": None}, ["Python"], "US") + assert result.points == [] + + def test_single_entry_unknown_granularity(self) -> None: + ts = int(datetime(2024, 1, 1).timestamp()) + data = {"timelineData": [{"time": str(ts), "value": "[50]"}]} + result = interest_over_time_to_result(data, ["Python"], "US") + assert result.granularity == "unknown" + assert len(result.points) == 1 + + def test_weekly_granularity(self, timeline_data_weekly: dict) -> None: + result = interest_over_time_to_result(timeline_data_weekly, ["Python", "JavaScript"], "US") + assert result.granularity == "weekly" + + def test_daily_granularity(self, timeline_data_daily: dict) -> None: + result = interest_over_time_to_result(timeline_data_daily, ["Python"], "US") + assert result.granularity == "daily" + + def test_hourly_granularity(self, timeline_data_hourly: dict) -> None: + result = interest_over_time_to_result(timeline_data_hourly, ["Python"], "US") + assert result.granularity == "hourly" + + def test_single_geo_scores_keyed_by_keyword(self, timeline_data_weekly: dict) -> None: + result = interest_over_time_to_result( + timeline_data_weekly, ["Python", "JavaScript"], "US" + ) + for point in result.points: + assert "Python" in point.scores + assert "JavaScript" in point.scores + + def test_single_geo_score_values(self, timeline_data_weekly: dict) -> None: + result = interest_over_time_to_result( + timeline_data_weekly, ["Python", "JavaScript"], "US" + ) + assert result.points[0].scores["Python"] == 80 + assert result.points[0].scores["JavaScript"] == 70 + + def test_multiple_geos_scores_keyed_with_pipe(self) -> None: + ts0 = int(datetime(2024, 1, 1).timestamp()) + ts1 = int(datetime(2024, 1, 8).timestamp()) + data = { + "timelineData": [ + {"time": str(ts0), "value": "[80, 60]"}, + {"time": str(ts1), "value": "[85, 65]"}, + ] + } + result = interest_over_time_to_result(data, ["Python"], ["US", "GB"]) + assert result.granularity == "weekly" + assert "Python|US" in result.points[0].scores + assert "Python|GB" in result.points[0].scores + + def test_multiple_geos_values(self) -> None: + ts0 = int(datetime(2024, 1, 1).timestamp()) + ts1 = int(datetime(2024, 1, 8).timestamp()) + data = { + "timelineData": [ + {"time": str(ts0), "value": "[80, 60]"}, + {"time": str(ts1), "value": "[85, 65]"}, + ] + } + result = interest_over_time_to_result(data, ["Python"], ["US", "GB"]) + assert result.points[0].scores["Python|US"] == 80 + assert result.points[0].scores["Python|GB"] == 60 + + def test_point_count_matches_timeline(self, timeline_data_weekly: dict) -> None: + result = interest_over_time_to_result( + timeline_data_weekly, ["Python", "JavaScript"], "US" + ) + assert len(result.points) == 3 + + def test_keywords_preserved(self, timeline_data_weekly: dict) -> None: + result = interest_over_time_to_result( + timeline_data_weekly, ["Python", "JavaScript"], "US" + ) + assert result.keywords == ["Python", "JavaScript"] + + def test_geo_as_list_with_single_element(self, timeline_data_daily: dict) -> None: + result = interest_over_time_to_result(timeline_data_daily, ["Python"], ["US"]) + assert result.granularity == "daily" + assert "Python" in result.points[0].scores + + def test_missing_value_field_treated_as_empty(self) -> None: + ts = int(datetime(2024, 1, 1).timestamp()) + data = {"timelineData": [{"time": str(ts)}]} + result = interest_over_time_to_result(data, ["Python"], "US") + assert len(result.points) == 1 + assert result.points[0].scores == {} + + +class TestInterestByRegionRows: + def test_basic_case(self) -> None: + data = { + "geoMapData": [ + {"geoName": "California", "value": "[90]"}, + {"geoName": "Texas", "value": "[70]"}, + ] + } + rows = interest_by_region_rows(data, "Python", ["Python"]) + assert len(rows) == 2 + assert rows[0].label == "California" + assert rows[0].value == 90 + assert rows[1].label == "Texas" + assert rows[1].value == 70 + + def test_empty_geo_map_data(self) -> None: + rows = interest_by_region_rows({"geoMapData": []}, "Python", ["Python"]) + assert rows == [] + + def test_none_geo_map_data(self) -> None: + rows = interest_by_region_rows({}, "Python", ["Python"]) + assert rows == [] + + def test_keyword_not_in_kw_list_uses_index_zero(self) -> None: + data = { + "geoMapData": [{"geoName": "UK", "value": "[50, 80]"}] + } + rows = interest_by_region_rows(data, "Unknown", ["Python", "JS"]) + assert rows[0].value == 50 + + def test_selects_correct_index_for_second_keyword(self) -> None: + data = { + "geoMapData": [{"geoName": "UK", "value": "[50, 80]"}] + } + rows = interest_by_region_rows(data, "JS", ["Python", "JS"]) + assert rows[0].value == 80 + + def test_value_index_out_of_range_returns_zero(self) -> None: + data = { + "geoMapData": [{"geoName": "UK", "value": "[50]"}] + } + rows = interest_by_region_rows(data, "JS", ["Python", "JS"]) + assert rows[0].value == 0 + + +class TestInterestByRegionToResult: + def test_returns_correct_type(self) -> None: + data = {"geoMapData": [{"geoName": "US", "value": "[100]"}]} + result = interest_by_region_to_result(data, "Python", ["Python"], Resolution.COUNTRY) + assert isinstance(result, InterestByRegionResult) + + def test_keyword_and_resolution_preserved(self) -> None: + data = {"geoMapData": [{"geoName": "US", "value": "[100]"}]} + result = interest_by_region_to_result(data, "Python", ["Python"], Resolution.REGION) + assert result.keyword == "Python" + assert result.resolution == Resolution.REGION + + +class TestTrendingTitlesToItems: + def test_basic_mapping(self) -> None: + items = trending_titles_to_items(["AI news", "Python 4"]) + assert len(items) == 2 + assert items[0].title == "AI news" + assert items[1].title == "Python 4" + + def test_empty_traffic_and_articles(self) -> None: + items = trending_titles_to_items(["test"]) + assert items[0].traffic == "" + assert items[0].articles == [] + + def test_empty_list(self) -> None: + assert trending_titles_to_items([]) == [] + + def test_non_string_titles_coerced(self) -> None: + items = trending_titles_to_items([42, None]) # type: ignore[list-item] + assert items[0].title == "42" + assert items[1].title == "None" + + +class TestToIntOrNone: + def test_none_returns_none(self) -> None: + assert _to_int_or_none(None) is None + + def test_nan_returns_none(self) -> None: + assert _to_int_or_none(float("nan")) is None + + def test_int_returns_int(self) -> None: + assert _to_int_or_none(42) == 42 + assert _to_int_or_none(0) == 0 + + def test_float_truncates(self) -> None: + assert _to_int_or_none(3.9) == 3 + + def test_bool_true(self) -> None: + assert _to_int_or_none(True) == 1 + + def test_bool_false(self) -> None: + assert _to_int_or_none(False) == 0 + + def test_string_int(self) -> None: + assert _to_int_or_none("100") == 100 + + def test_string_not_int_returns_none(self) -> None: + assert _to_int_or_none("not_a_number") is None + + def test_list_returns_none(self) -> None: + assert _to_int_or_none([1, 2]) is None + + +class TestParseTopRelated: + def test_none_returns_empty(self) -> None: + assert parse_top_related(None) == [] + + def test_empty_list_returns_empty(self) -> None: + assert parse_top_related([]) == [] + + def test_basic_row(self) -> None: + rows = [{"query": "python tutorial", "value": 100}] + result = parse_top_related(rows) + assert len(result) == 1 + assert result[0].term == "python tutorial" + assert result[0].value == 100 + + def test_missing_value_becomes_none(self) -> None: + rows = [{"query": "python", "value": float("nan")}] + result = parse_top_related(rows) + assert result[0].value is None + + def test_none_value_becomes_none(self) -> None: + rows = [{"query": "python", "value": None}] + result = parse_top_related(rows) + assert result[0].value is None + + def test_multiple_rows(self) -> None: + rows = [ + {"query": "python tutorial", "value": 100}, + {"query": "python course", "value": 80}, + ] + result = parse_top_related(rows) + assert len(result) == 2 + assert result[1].term == "python course" + assert result[1].value == 80 + + def test_missing_query_key_defaults_to_empty_string(self) -> None: + rows = [{"value": 50}] + result = parse_top_related(rows) + assert result[0].term == "" + + +class TestParseRisingRelated: + def test_none_returns_empty(self) -> None: + assert parse_rising_related(None) == [] + + def test_empty_list_returns_empty(self) -> None: + assert parse_rising_related([]) == [] + + def test_uses_formatted_value(self) -> None: + rows = [{"query": "python ai", "formattedValue": "+250%"}] + result = parse_rising_related(rows) + assert result[0].term == "python ai" + assert result[0].breakout == "+250%" + + def test_falls_back_to_value_if_no_formatted_value(self) -> None: + rows = [{"query": "python ai", "value": "Breakout"}] + result = parse_rising_related(rows) + assert result[0].breakout == "Breakout" + + def test_nan_breakout_becomes_none(self) -> None: + rows = [{"query": "python ai", "formattedValue": float("nan")}] + result = parse_rising_related(rows) + assert result[0].breakout is None + + def test_none_breakout_becomes_none(self) -> None: + rows = [{"query": "python ai", "formattedValue": None}] + result = parse_rising_related(rows) + assert result[0].breakout is None + + def test_breakout_coerced_to_string(self) -> None: + rows = [{"query": "test", "formattedValue": 300}] + result = parse_rising_related(rows) + assert result[0].breakout == "300" + + def test_value_field_never_set_on_rising(self) -> None: + rows = [{"query": "test", "formattedValue": "+100%"}] + result = parse_rising_related(rows) + assert result[0].value is None + + +class TestRelatedQueriesToResult: + def test_empty_raw_returns_empty_result(self) -> None: + result = related_queries_to_result({}, "Python") + assert result.top == [] + assert result.rising == [] + + def test_keyword_found(self) -> None: + raw = { + "Python": { + "top": [{"query": "python tutorial", "value": 100}], + "rising": [{"query": "python ai", "formattedValue": "+250%"}], + } + } + result = related_queries_to_result(raw, "Python") + assert len(result.top) == 1 + assert result.top[0].term == "python tutorial" + assert len(result.rising) == 1 + assert result.rising[0].breakout == "+250%" + + def test_keyword_not_found_single_bucket_fallback(self) -> None: + raw = { + "OtherKW": { + "top": [{"query": "something", "value": 50}], + "rising": [], + } + } + result = related_queries_to_result(raw, "Python") + assert len(result.top) == 1 + assert result.top[0].term == "something" + + def test_keyword_not_found_multiple_buckets_returns_empty(self) -> None: + raw = { + "A": {"top": [{"query": "a", "value": 1}], "rising": []}, + "B": {"top": [{"query": "b", "value": 2}], "rising": []}, + } + result = related_queries_to_result(raw, "Python") + assert result.top == [] + assert result.rising == [] + + def test_none_top_list(self) -> None: + raw = {"Python": {"top": None, "rising": []}} + result = related_queries_to_result(raw, "Python") + assert result.top == [] + + def test_none_rising_list(self) -> None: + raw = {"Python": {"top": [], "rising": None}} + result = related_queries_to_result(raw, "Python") + assert result.rising == [] + + def test_returns_related_result_type(self) -> None: + raw = {"Python": {"top": [], "rising": []}} + result = related_queries_to_result(raw, "Python") + assert isinstance(result, RelatedResult) + + +class TestTrendingResultFromTitles: + def test_returns_trending_result(self) -> None: + result = trending_result_from_titles(["AI", "Python"]) + assert isinstance(result, TrendingResult) + + def test_correct_item_count(self) -> None: + result = trending_result_from_titles(["A", "B", "C"]) + assert len(result.results) == 3 + + def test_titles_mapped(self) -> None: + result = trending_result_from_titles(["AI news"]) + assert result.results[0].title == "AI news" + + def test_empty_titles(self) -> None: + result = trending_result_from_titles([]) + assert result.results == [] diff --git a/tests/test_session.py b/tests/test_session.py new file mode 100644 index 0000000..d9929f7 --- /dev/null +++ b/tests/test_session.py @@ -0,0 +1,223 @@ +"""Tests for trendflow._trends_http.session helper functions and GoogleTrendsHttpSession.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from trendflow._trends_http.session import ( + GoogleTrendsHttpSession, + _normalize_proxies, + _primary_geo, +) +from trendflow._trends_http.transport import TrendsJsonTransport + + +def _make_session(**kwargs) -> GoogleTrendsHttpSession: + """Build a GoogleTrendsHttpSession with the transport patched out.""" + defaults: dict = {"hl": "en-US", "tz": 360} + defaults.update(kwargs) + with patch.object(TrendsJsonTransport, "_fetch_nid_cookies", return_value={}): + session = GoogleTrendsHttpSession(**defaults) + return session + + +class TestNormalizeProxies: + def test_empty_string_returns_empty_list(self) -> None: + assert _normalize_proxies("") == [] + + def test_non_empty_string_returns_single_element_list(self) -> None: + assert _normalize_proxies("http://proxy:8080") == ["http://proxy:8080"] + + def test_list_of_strings_returned_as_list(self) -> None: + proxies = ["http://p1:8080", "http://p2:8080"] + assert _normalize_proxies(proxies) == proxies + + def test_tuple_of_strings_returned_as_list(self) -> None: + result = _normalize_proxies(("http://p1:8080",)) + assert result == ["http://p1:8080"] + + def test_empty_list_returns_empty_list(self) -> None: + assert _normalize_proxies([]) == [] + + +class TestPrimaryGeo: + def test_string_returned_as_is(self) -> None: + assert _primary_geo("US") == "US" + + def test_empty_string_returned(self) -> None: + assert _primary_geo("") == "" + + def test_list_returns_first_element(self) -> None: + assert _primary_geo(["US", "GB"]) == "US" + + def test_empty_list_returns_empty_string(self) -> None: + assert _primary_geo([]) == "" + + def test_single_element_list(self) -> None: + assert _primary_geo(["DE"]) == "DE" + + +class TestGoogleTrendsHttpSessionInit: + def test_default_attributes(self) -> None: + session = _make_session() + assert session.hl == "en-US" + assert session.tz == 360 + assert session.kw_list == [] + assert session.token_payload == {} + + def test_geo_stored(self) -> None: + session = _make_session(geo="US") + assert session.geo == "US" + + def test_proxies_normalized(self) -> None: + session = _make_session(proxies="http://proxy:8080") + assert session.proxies == ["http://proxy:8080"] + + def test_empty_proxies(self) -> None: + session = _make_session(proxies="") + assert session.proxies == [] + + def test_cookies_property(self) -> None: + session = _make_session() + session._http.cookies = {"NID": "abc"} + assert session.cookies == {"NID": "abc"} + + def test_cookies_setter(self) -> None: + session = _make_session() + session.cookies = {"NID": "xyz"} + assert session._http.cookies == {"NID": "xyz"} + + def test_proxy_index_property(self) -> None: + session = _make_session() + assert session.proxy_index == 0 + + +class TestBuildPayload: + def test_invalid_gprop_raises_value_error(self) -> None: + session = _make_session() + with pytest.raises(ValueError, match="gprop"): + with patch.object(session, "_tokens"): + session.build_payload(["Python"], gprop="invalid") # type: ignore[arg-type] + + def test_valid_gprop_values(self) -> None: + session = _make_session() + for gprop in ("", "images", "news", "youtube", "froogle"): + with patch.object(session, "_tokens"): + session.build_payload(["Python"], gprop=gprop) # type: ignore[arg-type] + + def test_kw_list_stored(self) -> None: + session = _make_session() + with patch.object(session, "_tokens"): + session.build_payload(["Python", "JS"]) + assert session.kw_list == ["Python", "JS"] + + def test_geo_updated_when_provided(self) -> None: + session = _make_session(geo="") + with patch.object(session, "_tokens"): + session.build_payload(["Python"], geo="US") + assert "US" in session.geo + + def test_geo_preserved_when_not_provided(self) -> None: + session = _make_session(geo="DE") + with patch.object(session, "_tokens"): + session.build_payload(["Python"], geo="") + assert "DE" in session.geo + + def test_token_payload_has_req_key(self) -> None: + session = _make_session() + with patch.object(session, "_tokens"): + session.build_payload(["Python"]) + assert "req" in session.token_payload + + def test_token_payload_has_hl_and_tz(self) -> None: + session = _make_session(hl="en-US", tz=360) + with patch.object(session, "_tokens"): + session.build_payload(["Python"]) + assert session.token_payload["hl"] == "en-US" + assert session.token_payload["tz"] == 360 + + def test_calls_tokens(self) -> None: + session = _make_session() + with patch.object(session, "_tokens") as mock_tokens: + session.build_payload(["Python"]) + mock_tokens.assert_called_once() + + def test_list_timeframe_builds_per_item_payload(self) -> None: + session = _make_session() + with patch.object(session, "_tokens"): + session.build_payload(["Python"], timeframe=["today 12-m"]) + assert session.kw_list == ["Python"] + + def test_multiple_geos_creates_comparison_items(self) -> None: + session = _make_session() + with patch.object(session, "_tokens"): + session.build_payload(["Python"], geo="US") + + +class TestTopCharts: + def test_invalid_date_raises_value_error(self) -> None: + session = _make_session() + with pytest.raises(ValueError, match="year"): + session.top_charts("not-a-year") + + def test_none_date_raises_value_error(self) -> None: + session = _make_session() + with pytest.raises(ValueError): + session.top_charts(None) # type: ignore[arg-type] + + def test_valid_int_year(self) -> None: + session = _make_session() + mock_response = { + "topCharts": [{"listItems": [{"title": "item1"}]}] + } + with patch.object(session, "_get_data", return_value=mock_response): + result = session.top_charts(2023) + assert result == [{"title": "item1"}] + + def test_valid_string_year(self) -> None: + session = _make_session() + mock_response = { + "topCharts": [{"listItems": [{"title": "item1"}]}] + } + with patch.object(session, "_get_data", return_value=mock_response): + result = session.top_charts("2023") + assert result == [{"title": "item1"}] + + def test_empty_top_charts_returns_none(self) -> None: + session = _make_session() + mock_response = {"topCharts": []} + with patch.object(session, "_get_data", return_value=mock_response): + result = session.top_charts(2023) + assert result is None + + +class TestTrendingSearches: + def test_returns_list_for_pn(self) -> None: + session = _make_session() + mock_response = {"united_states": ["AI", "Python"]} + with patch.object(session, "_get_data", return_value=mock_response): + result = session.trending_searches(pn="united_states") + assert result == ["AI", "Python"] + + def test_returns_list_type(self) -> None: + session = _make_session() + mock_response = {"germany": ["Bayern", "Bundesliga"]} + with patch.object(session, "_get_data", return_value=mock_response): + result = session.trending_searches(pn="germany") + assert isinstance(result, list) + + +class TestRelatedQueriesWidgets: + def test_empty_widget_list_returns_empty_dict(self) -> None: + session = _make_session() + session.related_queries_widget_list = [] + result = session.related_queries() + assert result == {} + + def test_empty_topics_widget_list_returns_empty_dict(self) -> None: + session = _make_session() + session.related_topics_widget_list = [] + result = session.related_topics() + assert result == {} diff --git a/tests/test_transport.py b/tests/test_transport.py new file mode 100644 index 0000000..520fee2 --- /dev/null +++ b/tests/test_transport.py @@ -0,0 +1,270 @@ +"""Tests for trendflow._trends_http.transport helper functions and TrendsJsonTransport.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from trendflow._trends_http.exceptions import ResponseError, TooManyRequestsError +from trendflow._trends_http.transport import ( + TrendsJsonTransport, + _extra_for_httpx, + _json_content_type, + _normalize_timeout, +) + + +class TestNormalizeTimeout: + def test_httpx_timeout_passthrough(self) -> None: + t = httpx.Timeout(10.0) + assert _normalize_timeout(t) is t + + def test_float_becomes_httpx_timeout(self) -> None: + result = _normalize_timeout(5.0) + assert isinstance(result, httpx.Timeout) + + def test_int_becomes_httpx_timeout(self) -> None: + result = _normalize_timeout(10) + assert isinstance(result, httpx.Timeout) + + def test_tuple_sets_connect_and_read(self) -> None: + result = _normalize_timeout((2.0, 10.0)) + assert isinstance(result, httpx.Timeout) + assert result.connect == 2.0 + assert result.read == 10.0 + + def test_tuple_values_coerced_to_float(self) -> None: + result = _normalize_timeout((2, 10)) + assert result.connect == 2.0 + assert result.read == 10.0 + + +class TestExtraForHttpx: + def test_passes_through_non_proxies_keys(self) -> None: + extra = {"verify": False, "follow_redirects": True} + result = _extra_for_httpx(extra) + assert result["verify"] is False + assert result["follow_redirects"] is True + + def test_removes_proxies_key(self) -> None: + extra = {"proxies": {"https://": "http://proxy:8080"}} + result = _extra_for_httpx(extra) + assert "proxies" not in result + + def test_maps_proxies_dict_to_proxy(self) -> None: + extra = {"proxies": {"https://": "http://proxy:8080"}} + result = _extra_for_httpx(extra) + assert result["proxy"] == "http://proxy:8080" + + def test_maps_proxies_string_to_proxy(self) -> None: + extra = {"proxies": "http://proxy:8080"} + result = _extra_for_httpx(extra) + assert result["proxy"] == "http://proxy:8080" + + def test_none_proxies_not_added(self) -> None: + extra = {"proxies": None} + result = _extra_for_httpx(extra) + assert "proxy" not in result + assert "proxies" not in result + + def test_empty_dict_stays_empty(self) -> None: + result = _extra_for_httpx({}) + assert result == {} + + +class TestJsonContentType: + def test_application_json(self) -> None: + assert _json_content_type("application/json") is True + + def test_application_json_with_charset(self) -> None: + assert _json_content_type("application/json; charset=utf-8") is True + + def test_application_javascript(self) -> None: + assert _json_content_type("application/javascript") is True + + def test_text_javascript(self) -> None: + assert _json_content_type("text/javascript") is True + + def test_text_html_is_false(self) -> None: + assert _json_content_type("text/html") is False + + def test_text_plain_is_false(self) -> None: + assert _json_content_type("text/plain") is False + + def test_empty_string_is_false(self) -> None: + assert _json_content_type("") is False + + def test_case_insensitive(self) -> None: + assert _json_content_type("Application/JSON") is True + + +def _make_transport(**kwargs) -> TrendsJsonTransport: + """Build a TrendsJsonTransport with _fetch_nid_cookies patched out.""" + defaults = dict( + hl="en-US", + tz=360, + timeout=(2.0, 5.0), + headers={}, + extra_client_args={}, + proxy_urls=[], + retries=0, + ) + defaults.update(kwargs) + with patch.object(TrendsJsonTransport, "_fetch_nid_cookies", return_value={"NID": "test"}): + transport = TrendsJsonTransport(**defaults) + return transport + + +class TestTrendsJsonTransportAdvanceProxy: + def test_no_proxies_is_noop(self) -> None: + transport = _make_transport(proxy_urls=[]) + transport.advance_proxy() + assert transport.proxy_index == 0 + + def test_single_proxy_wraps_to_zero(self) -> None: + transport = _make_transport(proxy_urls=["http://proxy1:8080"]) + assert transport.proxy_index == 0 + transport.advance_proxy() + assert transport.proxy_index == 0 + + def test_two_proxies_advance(self) -> None: + transport = _make_transport(proxy_urls=["http://p1:8080", "http://p2:8080"]) + assert transport.proxy_index == 0 + transport.advance_proxy() + assert transport.proxy_index == 1 + + def test_two_proxies_wraps_after_last(self) -> None: + transport = _make_transport(proxy_urls=["http://p1:8080", "http://p2:8080"]) + transport.advance_proxy() # → 1 + transport.advance_proxy() # → 0 (wrap) + assert transport.proxy_index == 0 + + +class TestTrendsJsonTransportCookieUrl: + def test_cookie_url_uses_last_two_chars_of_hl(self) -> None: + transport = _make_transport(hl="en-US") + url = transport._explore_cookie_url() + assert url.endswith("?geo=US") + + def test_cookie_url_for_gb(self) -> None: + transport = _make_transport(hl="en-GB") + url = transport._explore_cookie_url() + assert url.endswith("?geo=GB") + + +class TestTrendsJsonTransportRequestJson: + def _mock_response(self, status_code: int, content_type: str, body: str) -> MagicMock: + resp = MagicMock(spec=httpx.Response) + resp.status_code = status_code + resp.headers = {"content-type": content_type} + resp.text = body + return resp + + def test_get_returns_parsed_json(self) -> None: + transport = _make_transport() + payload = {"key": "value"} + response = self._mock_response(200, "application/json", json.dumps(payload)) + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + result = transport.request_json("https://example.com", "get") + + assert result == payload + + def test_post_calls_client_post(self) -> None: + transport = _make_transport() + payload = {"data": [1, 2, 3]} + response = self._mock_response(200, "application/json", json.dumps(payload)) + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.post.return_value = response + mock_client_cls.return_value = mock_client + + result = transport.request_json("https://example.com", "post") + + assert result == payload + mock_client.post.assert_called_once() + + def test_trim_chars_strips_prefix(self) -> None: + transport = _make_transport() + response = self._mock_response(200, "application/json", ")]}'\n{\"x\": 1}") + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + result = transport.request_json("https://example.com", "get", trim_chars=5) + + assert result == {"x": 1} + + def test_429_raises_too_many_requests(self) -> None: + transport = _make_transport() + response = self._mock_response(429, "application/json", "{}") + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + with pytest.raises(TooManyRequestsError): + transport.request_json("https://example.com", "get") + + def test_500_raises_response_error(self) -> None: + transport = _make_transport() + response = self._mock_response(500, "text/html", "error") + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + with pytest.raises(ResponseError): + transport.request_json("https://example.com", "get") + + def test_non_json_200_raises_response_error(self) -> None: + transport = _make_transport() + response = self._mock_response(200, "text/html", "") + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + with pytest.raises(ResponseError): + transport.request_json("https://example.com", "get") + + def test_successful_request_advances_proxy(self) -> None: + transport = _make_transport(proxy_urls=["http://p1:8080", "http://p2:8080"]) + response = self._mock_response(200, "application/json", "{}") + + with patch("httpx.Client") as mock_client_cls: + mock_client = MagicMock() + mock_client.__enter__ = MagicMock(return_value=mock_client) + mock_client.__exit__ = MagicMock(return_value=False) + mock_client.get.return_value = response + mock_client_cls.return_value = mock_client + + with patch.object(transport, "_fetch_nid_cookies", return_value={}): + transport.request_json("https://example.com", "get") + + assert transport.proxy_index == 1 diff --git a/tests/test_trendflow.py b/tests/test_trendflow.py index 3e3107f..349d967 100644 --- a/tests/test_trendflow.py +++ b/tests/test_trendflow.py @@ -1,8 +1,55 @@ -"""Tests for `trendflow` package.""" +"""Tests for `trendflow` package — public API surface.""" import trendflow +from trendflow import ( + Client, + ExportFormat, + GoogleTrendsFetcher, + InterestByRegionResult, + InterestOverTimeResult, + Region, + RelatedQuery, + RelatedResult, + Resolution, + Timeframe, + TrendingItem, + TrendingResult, + TrendPoint, + TrendsFetcher, +) +from trendflow.models import RegionalInterestRow def test_import(): """Verify the package can be imported.""" assert trendflow + + +def test_client_is_alias_for_google_trends_fetcher(): + assert Client is GoogleTrendsFetcher + + +def test_all_enums_exported(): + assert Region is not None + assert Timeframe is not None + assert Resolution is not None + assert ExportFormat is not None + + +def test_all_models_exported(): + assert TrendPoint is not None + assert InterestOverTimeResult is not None + assert RegionalInterestRow is not None + assert InterestByRegionResult is not None + assert TrendingItem is not None + assert TrendingResult is not None + assert RelatedQuery is not None + assert RelatedResult is not None + + +def test_trends_fetcher_protocol_exported(): + assert TrendsFetcher is not None + + +def test_google_trends_fetcher_exported(): + assert GoogleTrendsFetcher is not None From 80c0c9f841beace68ce5020b016d020eab715d78 Mon Sep 17 00:00:00 2001 From: Vincent WENDLING Date: Sat, 18 Apr 2026 23:56:37 +0200 Subject: [PATCH 2/3] fix linting and type checking --- tests/conftest.py | 4 +- tests/test_exporters.py | 2 +- tests/test_fetcher.py | 90 ++++++++++++++++++++--------------------- tests/test_models.py | 20 ++++----- tests/test_parsers.py | 39 ++++++------------ tests/test_session.py | 14 +++---- tests/test_transport.py | 35 ++++++++++------ 7 files changed, 93 insertions(+), 111 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index 527dc68..a1e8afe 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -62,9 +62,7 @@ def region_rows() -> list[RegionalInterestRow]: @pytest.fixture def ibr_result(region_rows: list[RegionalInterestRow]) -> InterestByRegionResult: - return InterestByRegionResult( - keyword="Python", resolution=Resolution.REGION, rows=region_rows - ) + return InterestByRegionResult(keyword="Python", resolution=Resolution.REGION, rows=region_rows) @pytest.fixture diff --git a/tests/test_exporters.py b/tests/test_exporters.py index 459a618..218da6a 100644 --- a/tests/test_exporters.py +++ b/tests/test_exporters.py @@ -179,7 +179,7 @@ def test_dispatches_json(self, result_with_points: InterestOverTimeResult) -> No def test_unsupported_format_raises_value_error(self, result_with_points: InterestOverTimeResult) -> None: with pytest.raises(ValueError, match="Unsupported export format"): - export_interest_over_time(result_with_points, "xml", Path("/tmp/out.xml")) # type: ignore[arg-type] + export_interest_over_time(result_with_points, "xml", Path("/tmp/out.xml")) # type: ignore def test_exporters_registry_has_csv_and_json(self) -> None: assert ExportFormat.CSV in INTEREST_OVER_TIME_EXPORTERS diff --git a/tests/test_fetcher.py b/tests/test_fetcher.py index 4611b66..8d80a2a 100644 --- a/tests/test_fetcher.py +++ b/tests/test_fetcher.py @@ -2,6 +2,7 @@ from __future__ import annotations +from typing import Any from unittest.mock import MagicMock, patch import pytest @@ -51,14 +52,14 @@ def test_all_non_worldwide_regions_have_mapping(self) -> None: assert region in TRENDING_PN, f"Missing TRENDING_PN entry for {region!r}" -def _make_fetcher() -> GoogleTrendsFetcher: - """Build a GoogleTrendsFetcher with GoogleTrendsHttpSession patched out.""" +def _make_fetcher() -> tuple[GoogleTrendsFetcher, Any]: + """Return (fetcher, mock_req) with GoogleTrendsHttpSession patched out.""" with patch("trendflow._fetcher.GoogleTrendsHttpSession") as mock_session_cls: mock_session = MagicMock() mock_session.geo = "US" mock_session_cls.return_value = mock_session fetcher = GoogleTrendsFetcher() - return fetcher + return fetcher, mock_session class TestGoogleTrendsFetcherInit: @@ -97,11 +98,11 @@ def test_timeout_min_read_is_timeout_plus_5(self) -> None: class TestTrendsFetcherProtocol: def test_google_trends_fetcher_implements_protocol(self) -> None: - fetcher = _make_fetcher() + fetcher, _ = _make_fetcher() assert isinstance(fetcher, TrendsFetcher) def test_protocol_methods_exist(self) -> None: - fetcher = _make_fetcher() + fetcher, _ = _make_fetcher() assert hasattr(fetcher, "interest_over_time") assert hasattr(fetcher, "interest_by_region") assert hasattr(fetcher, "trending_now") @@ -110,10 +111,9 @@ def test_protocol_methods_exist(self) -> None: class TestInterestOverTime: def test_calls_build_payload_with_correct_args(self) -> None: - fetcher = _make_fetcher() - mock_default = {"timelineData": []} - fetcher._req.interest_over_time.return_value = mock_default - fetcher._req.geo = "US" + fetcher, req = _make_fetcher() + req.interest_over_time.return_value = {"timelineData": []} + req.geo = "US" fetcher.interest_over_time( keywords=["Python"], @@ -121,7 +121,7 @@ def test_calls_build_payload_with_correct_args(self) -> None: region=Region.US, ) - fetcher._req.build_payload.assert_called_once_with( + req.build_payload.assert_called_once_with( ["Python"], cat=0, timeframe=Timeframe.PAST_YEAR.value, @@ -130,9 +130,9 @@ def test_calls_build_payload_with_correct_args(self) -> None: ) def test_returns_interest_over_time_result(self) -> None: - fetcher = _make_fetcher() - fetcher._req.interest_over_time.return_value = {"timelineData": []} - fetcher._req.geo = "US" + fetcher, req = _make_fetcher() + req.interest_over_time.return_value = {"timelineData": []} + req.geo = "US" result = fetcher.interest_over_time( keywords=["Python"], @@ -143,9 +143,9 @@ def test_returns_interest_over_time_result(self) -> None: assert isinstance(result, InterestOverTimeResult) def test_passes_geo_from_session(self) -> None: - fetcher = _make_fetcher() - fetcher._req.interest_over_time.return_value = {"timelineData": []} - fetcher._req.geo = ["US"] + fetcher, req = _make_fetcher() + req.interest_over_time.return_value = {"timelineData": []} + req.geo = ["US"] result = fetcher.interest_over_time( keywords=["Python"], @@ -158,8 +158,8 @@ def test_passes_geo_from_session(self) -> None: class TestInterestByRegion: def test_returns_empty_result_when_no_geo_map_data(self) -> None: - fetcher = _make_fetcher() - fetcher._req.interest_by_region.return_value = {} + fetcher, req = _make_fetcher() + req.interest_by_region.return_value = {} result = fetcher.interest_by_region( keyword="Python", @@ -172,12 +172,12 @@ def test_returns_empty_result_when_no_geo_map_data(self) -> None: assert result.keyword == "Python" def test_calls_build_payload_with_keyword(self) -> None: - fetcher = _make_fetcher() - fetcher._req.interest_by_region.return_value = {} + fetcher, req = _make_fetcher() + req.interest_by_region.return_value = {} fetcher.interest_by_region(keyword="Rust", resolution=Resolution.REGION, region=Region.US) - fetcher._req.build_payload.assert_called_once_with( + req.build_payload.assert_called_once_with( ["Rust"], cat=0, timeframe=Timeframe.PAST_YEAR.value, @@ -186,14 +186,10 @@ def test_calls_build_payload_with_keyword(self) -> None: ) def test_returns_parsed_result_when_data_present(self) -> None: - fetcher = _make_fetcher() - fetcher._req.interest_by_region.return_value = { - "geoMapData": [{"geoName": "California", "value": "[90]"}] - } + fetcher, req = _make_fetcher() + req.interest_by_region.return_value = {"geoMapData": [{"geoName": "California", "value": "[90]"}]} - result = fetcher.interest_by_region( - keyword="Python", resolution=Resolution.REGION, region=Region.US - ) + result = fetcher.interest_by_region(keyword="Python", resolution=Resolution.REGION, region=Region.US) assert isinstance(result, InterestByRegionResult) assert len(result.rows) == 1 @@ -202,22 +198,22 @@ def test_returns_parsed_result_when_data_present(self) -> None: class TestTrendingNow: def test_worldwide_raises_value_error(self) -> None: - fetcher = _make_fetcher() + fetcher, _ = _make_fetcher() with pytest.raises(ValueError, match="specific country"): fetcher.trending_now(region=Region.WORLDWIDE) def test_valid_region_calls_trending_searches(self) -> None: - fetcher = _make_fetcher() - fetcher._req.trending_searches.return_value = ["AI", "Python"] + fetcher, req = _make_fetcher() + req.trending_searches.return_value = ["AI", "Python"] result = fetcher.trending_now(region=Region.US) - fetcher._req.trending_searches.assert_called_once_with(pn="united_states") + req.trending_searches.assert_called_once_with(pn="united_states") assert isinstance(result, TrendingResult) def test_returns_correct_titles(self) -> None: - fetcher = _make_fetcher() - fetcher._req.trending_searches.return_value = ["AI tools", "Python 4"] + fetcher, req = _make_fetcher() + req.trending_searches.return_value = ["AI tools", "Python 4"] result = fetcher.trending_now(region=Region.US) @@ -225,30 +221,30 @@ def test_returns_correct_titles(self) -> None: assert result.results[1].title == "Python 4" def test_pn_lookup_for_gb(self) -> None: - fetcher = _make_fetcher() - fetcher._req.trending_searches.return_value = [] + fetcher, req = _make_fetcher() + req.trending_searches.return_value = [] fetcher.trending_now(region=Region.GB) - fetcher._req.trending_searches.assert_called_once_with(pn="united_kingdom") + req.trending_searches.assert_called_once_with(pn="united_kingdom") def test_pn_lookup_for_de(self) -> None: - fetcher = _make_fetcher() - fetcher._req.trending_searches.return_value = [] + fetcher, req = _make_fetcher() + req.trending_searches.return_value = [] fetcher.trending_now(region=Region.DE) - fetcher._req.trending_searches.assert_called_once_with(pn="germany") + req.trending_searches.assert_called_once_with(pn="germany") class TestRelatedQueries: def test_calls_build_payload_with_keyword(self) -> None: - fetcher = _make_fetcher() - fetcher._req.related_queries.return_value = {} + fetcher, req = _make_fetcher() + req.related_queries.return_value = {} fetcher.related_queries(keyword="Python") - fetcher._req.build_payload.assert_called_once_with( + req.build_payload.assert_called_once_with( ["Python"], cat=0, timeframe=Timeframe.PAST_YEAR.value, @@ -257,8 +253,8 @@ def test_calls_build_payload_with_keyword(self) -> None: ) def test_returns_related_result(self) -> None: - fetcher = _make_fetcher() - fetcher._req.related_queries.return_value = { + fetcher, req = _make_fetcher() + req.related_queries.return_value = { "Python": { "top": [{"query": "python tutorial", "value": 100}], "rising": [], @@ -272,8 +268,8 @@ def test_returns_related_result(self) -> None: assert result.top[0].term == "python tutorial" def test_empty_raw_returns_empty_result(self) -> None: - fetcher = _make_fetcher() - fetcher._req.related_queries.return_value = {} + fetcher, req = _make_fetcher() + req.related_queries.return_value = {} result = fetcher.related_queries(keyword="Python") diff --git a/tests/test_models.py b/tests/test_models.py index f2f321d..9e37c52 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -33,7 +33,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: point = TrendPoint(date=datetime(2024, 1, 1), scores={"Python": 80}) with pytest.raises(Exception): - point.date = datetime(2024, 1, 2) # type: ignore[misc] + point.date = datetime(2024, 1, 2) # type: ignore def test_equality(self) -> None: dt = datetime(2024, 1, 1) @@ -59,9 +59,7 @@ def test_to_dataframe_empty(self, empty_iot_result: InterestOverTimeResult) -> N assert "Python" in df.columns def test_to_dataframe_empty_has_keyword_columns(self) -> None: - result = InterestOverTimeResult( - keywords=["A", "B", "C"], granularity="unknown", points=[] - ) + result = InterestOverTimeResult(keywords=["A", "B", "C"], granularity="unknown", points=[]) df = result.to_dataframe() assert list(df.columns) == ["date", "A", "B", "C"] @@ -84,7 +82,7 @@ def test_to_dataframe_date_column(self, iot_result: InterestOverTimeResult) -> N def test_frozen(self, iot_result: InterestOverTimeResult) -> None: with pytest.raises(Exception): - iot_result.keywords = ["other"] # type: ignore[misc] + iot_result.keywords = ["other"] # type: ignore def test_export_csv(self, iot_result: InterestOverTimeResult) -> None: with tempfile.NamedTemporaryFile(suffix=".csv", delete=False) as f: @@ -138,7 +136,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: row = RegionalInterestRow(label="California", value=90) with pytest.raises(Exception): - row.value = 100 # type: ignore[misc] + row.value = 100 # type: ignore class TestInterestByRegionResult: @@ -149,12 +147,10 @@ def test_construction(self, ibr_result: InterestByRegionResult) -> None: def test_frozen(self, ibr_result: InterestByRegionResult) -> None: with pytest.raises(Exception): - ibr_result.keyword = "other" # type: ignore[misc] + ibr_result.keyword = "other" # type: ignore def test_empty_rows(self) -> None: - result = InterestByRegionResult( - keyword="Python", resolution=Resolution.COUNTRY, rows=[] - ) + result = InterestByRegionResult(keyword="Python", resolution=Resolution.COUNTRY, rows=[]) assert result.rows == [] @@ -168,7 +164,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: item = TrendingItem(title="AI news", traffic="500K+", articles=[]) with pytest.raises(Exception): - item.title = "other" # type: ignore[misc] + item.title = "other" # type: ignore class TestTrendingResult: @@ -198,7 +194,7 @@ def test_with_breakout(self) -> None: def test_frozen(self) -> None: q = RelatedQuery(term="python tutorial", value=100) with pytest.raises(Exception): - q.term = "other" # type: ignore[misc] + q.term = "other" # type: ignore class TestRelatedResult: diff --git a/tests/test_parsers.py b/tests/test_parsers.py index 8b92c89..73d3c4a 100644 --- a/tests/test_parsers.py +++ b/tests/test_parsers.py @@ -4,6 +4,7 @@ import math from datetime import datetime, timedelta +from typing import Any import pytest @@ -162,17 +163,13 @@ def test_hourly_granularity(self, timeline_data_hourly: dict) -> None: assert result.granularity == "hourly" def test_single_geo_scores_keyed_by_keyword(self, timeline_data_weekly: dict) -> None: - result = interest_over_time_to_result( - timeline_data_weekly, ["Python", "JavaScript"], "US" - ) + result = interest_over_time_to_result(timeline_data_weekly, ["Python", "JavaScript"], "US") for point in result.points: assert "Python" in point.scores assert "JavaScript" in point.scores def test_single_geo_score_values(self, timeline_data_weekly: dict) -> None: - result = interest_over_time_to_result( - timeline_data_weekly, ["Python", "JavaScript"], "US" - ) + result = interest_over_time_to_result(timeline_data_weekly, ["Python", "JavaScript"], "US") assert result.points[0].scores["Python"] == 80 assert result.points[0].scores["JavaScript"] == 70 @@ -204,15 +201,11 @@ def test_multiple_geos_values(self) -> None: assert result.points[0].scores["Python|GB"] == 60 def test_point_count_matches_timeline(self, timeline_data_weekly: dict) -> None: - result = interest_over_time_to_result( - timeline_data_weekly, ["Python", "JavaScript"], "US" - ) + result = interest_over_time_to_result(timeline_data_weekly, ["Python", "JavaScript"], "US") assert len(result.points) == 3 def test_keywords_preserved(self, timeline_data_weekly: dict) -> None: - result = interest_over_time_to_result( - timeline_data_weekly, ["Python", "JavaScript"], "US" - ) + result = interest_over_time_to_result(timeline_data_weekly, ["Python", "JavaScript"], "US") assert result.keywords == ["Python", "JavaScript"] def test_geo_as_list_with_single_element(self, timeline_data_daily: dict) -> None: @@ -252,23 +245,17 @@ def test_none_geo_map_data(self) -> None: assert rows == [] def test_keyword_not_in_kw_list_uses_index_zero(self) -> None: - data = { - "geoMapData": [{"geoName": "UK", "value": "[50, 80]"}] - } + data = {"geoMapData": [{"geoName": "UK", "value": "[50, 80]"}]} rows = interest_by_region_rows(data, "Unknown", ["Python", "JS"]) assert rows[0].value == 50 def test_selects_correct_index_for_second_keyword(self) -> None: - data = { - "geoMapData": [{"geoName": "UK", "value": "[50, 80]"}] - } + data = {"geoMapData": [{"geoName": "UK", "value": "[50, 80]"}]} rows = interest_by_region_rows(data, "JS", ["Python", "JS"]) assert rows[0].value == 80 def test_value_index_out_of_range_returns_zero(self) -> None: - data = { - "geoMapData": [{"geoName": "UK", "value": "[50]"}] - } + data = {"geoMapData": [{"geoName": "UK", "value": "[50]"}]} rows = interest_by_region_rows(data, "JS", ["Python", "JS"]) assert rows[0].value == 0 @@ -302,7 +289,7 @@ def test_empty_list(self) -> None: assert trending_titles_to_items([]) == [] def test_non_string_titles_coerced(self) -> None: - items = trending_titles_to_items([42, None]) # type: ignore[list-item] + items = trending_titles_to_items([42, None]) # type: ignore assert items[0].title == "42" assert items[1].title == "None" @@ -423,7 +410,7 @@ def test_empty_raw_returns_empty_result(self) -> None: assert result.rising == [] def test_keyword_found(self) -> None: - raw = { + raw: dict[str, dict[str, list[dict[str, Any]] | None]] = { "Python": { "top": [{"query": "python tutorial", "value": 100}], "rising": [{"query": "python ai", "formattedValue": "+250%"}], @@ -436,7 +423,7 @@ def test_keyword_found(self) -> None: assert result.rising[0].breakout == "+250%" def test_keyword_not_found_single_bucket_fallback(self) -> None: - raw = { + raw: dict[str, dict[str, list[dict[str, Any]] | None]] = { "OtherKW": { "top": [{"query": "something", "value": 50}], "rising": [], @@ -447,7 +434,7 @@ def test_keyword_not_found_single_bucket_fallback(self) -> None: assert result.top[0].term == "something" def test_keyword_not_found_multiple_buckets_returns_empty(self) -> None: - raw = { + raw: dict[str, dict[str, list[dict[str, Any]] | None]] = { "A": {"top": [{"query": "a", "value": 1}], "rising": []}, "B": {"top": [{"query": "b", "value": 2}], "rising": []}, } @@ -466,7 +453,7 @@ def test_none_rising_list(self) -> None: assert result.rising == [] def test_returns_related_result_type(self) -> None: - raw = {"Python": {"top": [], "rising": []}} + raw: dict[str, dict[str, list[dict[str, Any]] | None]] = {"Python": {"top": [], "rising": []}} result = related_queries_to_result(raw, "Python") assert isinstance(result, RelatedResult) diff --git a/tests/test_session.py b/tests/test_session.py index d9929f7..af33007 100644 --- a/tests/test_session.py +++ b/tests/test_session.py @@ -99,13 +99,13 @@ def test_invalid_gprop_raises_value_error(self) -> None: session = _make_session() with pytest.raises(ValueError, match="gprop"): with patch.object(session, "_tokens"): - session.build_payload(["Python"], gprop="invalid") # type: ignore[arg-type] + session.build_payload(["Python"], gprop="invalid") # type: ignore def test_valid_gprop_values(self) -> None: session = _make_session() for gprop in ("", "images", "news", "youtube", "froogle"): with patch.object(session, "_tokens"): - session.build_payload(["Python"], gprop=gprop) # type: ignore[arg-type] + session.build_payload(["Python"], gprop=gprop) def test_kw_list_stored(self) -> None: session = _make_session() @@ -165,22 +165,18 @@ def test_invalid_date_raises_value_error(self) -> None: def test_none_date_raises_value_error(self) -> None: session = _make_session() with pytest.raises(ValueError): - session.top_charts(None) # type: ignore[arg-type] + session.top_charts(None) # type: ignore def test_valid_int_year(self) -> None: session = _make_session() - mock_response = { - "topCharts": [{"listItems": [{"title": "item1"}]}] - } + mock_response = {"topCharts": [{"listItems": [{"title": "item1"}]}]} with patch.object(session, "_get_data", return_value=mock_response): result = session.top_charts(2023) assert result == [{"title": "item1"}] def test_valid_string_year(self) -> None: session = _make_session() - mock_response = { - "topCharts": [{"listItems": [{"title": "item1"}]}] - } + mock_response = {"topCharts": [{"listItems": [{"title": "item1"}]}]} with patch.object(session, "_get_data", return_value=mock_response): result = session.top_charts("2023") assert result == [{"title": "item1"}] diff --git a/tests/test_transport.py b/tests/test_transport.py index 520fee2..a63150a 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -3,6 +3,8 @@ from __future__ import annotations import json +from collections.abc import Mapping, MutableMapping +from typing import Any from unittest.mock import MagicMock, patch import httpx @@ -101,20 +103,27 @@ def test_case_insensitive(self) -> None: assert _json_content_type("Application/JSON") is True -def _make_transport(**kwargs) -> TrendsJsonTransport: +def _make_transport( + *, + hl: str = "en-US", + tz: int = 360, + timeout: httpx.Timeout | tuple[float, float] | float = (2.0, 5.0), + headers: MutableMapping[str, str] | None = None, + extra_client_args: Mapping[str, Any] | None = None, + proxy_urls: list[str] | None = None, + retries: int = 0, +) -> TrendsJsonTransport: """Build a TrendsJsonTransport with _fetch_nid_cookies patched out.""" - defaults = dict( - hl="en-US", - tz=360, - timeout=(2.0, 5.0), - headers={}, - extra_client_args={}, - proxy_urls=[], - retries=0, - ) - defaults.update(kwargs) with patch.object(TrendsJsonTransport, "_fetch_nid_cookies", return_value={"NID": "test"}): - transport = TrendsJsonTransport(**defaults) + transport = TrendsJsonTransport( + hl=hl, + tz=tz, + timeout=timeout, + headers=headers if headers is not None else {}, + extra_client_args=extra_client_args if extra_client_args is not None else {}, + proxy_urls=proxy_urls if proxy_urls is not None else [], + retries=retries, + ) return transport @@ -198,7 +207,7 @@ def test_post_calls_client_post(self) -> None: def test_trim_chars_strips_prefix(self) -> None: transport = _make_transport() - response = self._mock_response(200, "application/json", ")]}'\n{\"x\": 1}") + response = self._mock_response(200, "application/json", ')]}\'\n{"x": 1}') with patch("httpx.Client") as mock_client_cls: mock_client = MagicMock() From f7d143b96f69c94fc31739aa755f5efae57a20ca Mon Sep 17 00:00:00 2001 From: Vincent WENDLING Date: Sun, 19 Apr 2026 00:01:08 +0200 Subject: [PATCH 3/3] chore: fix linting error (forgot to commit some changes) --- tests/conftest.py | 2 +- tests/test_enums.py | 2 -- tests/test_models.py | 13 +++++++------ tests/test_parsers.py | 5 ----- tests/test_session.py | 2 +- 5 files changed, 9 insertions(+), 15 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index a1e8afe..2614814 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -6,7 +6,7 @@ import pytest -from trendflow.enums import ExportFormat, Region, Resolution, Timeframe +from trendflow.enums import Resolution from trendflow.models import ( InterestByRegionResult, InterestOverTimeResult, diff --git a/tests/test_enums.py b/tests/test_enums.py index 5209d83..d57acdd 100644 --- a/tests/test_enums.py +++ b/tests/test_enums.py @@ -2,8 +2,6 @@ from __future__ import annotations -import pytest - from trendflow.enums import ExportFormat, Region, Resolution, Timeframe diff --git a/tests/test_models.py b/tests/test_models.py index 9e37c52..ae9af9e 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -2,6 +2,7 @@ from __future__ import annotations +import dataclasses import json import tempfile from datetime import datetime @@ -32,7 +33,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: point = TrendPoint(date=datetime(2024, 1, 1), scores={"Python": 80}) - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): point.date = datetime(2024, 1, 2) # type: ignore def test_equality(self) -> None: @@ -81,7 +82,7 @@ def test_to_dataframe_date_column(self, iot_result: InterestOverTimeResult) -> N assert df["date"].iloc[0] == datetime(2024, 1, 1) def test_frozen(self, iot_result: InterestOverTimeResult) -> None: - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): iot_result.keywords = ["other"] # type: ignore def test_export_csv(self, iot_result: InterestOverTimeResult) -> None: @@ -135,7 +136,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: row = RegionalInterestRow(label="California", value=90) - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): row.value = 100 # type: ignore @@ -146,7 +147,7 @@ def test_construction(self, ibr_result: InterestByRegionResult) -> None: assert len(ibr_result.rows) == 2 def test_frozen(self, ibr_result: InterestByRegionResult) -> None: - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): ibr_result.keyword = "other" # type: ignore def test_empty_rows(self) -> None: @@ -163,7 +164,7 @@ def test_construction(self) -> None: def test_frozen(self) -> None: item = TrendingItem(title="AI news", traffic="500K+", articles=[]) - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): item.title = "other" # type: ignore @@ -193,7 +194,7 @@ def test_with_breakout(self) -> None: def test_frozen(self) -> None: q = RelatedQuery(term="python tutorial", value=100) - with pytest.raises(Exception): + with pytest.raises(dataclasses.FrozenInstanceError): q.term = "other" # type: ignore diff --git a/tests/test_parsers.py b/tests/test_parsers.py index 73d3c4a..64ecf76 100644 --- a/tests/test_parsers.py +++ b/tests/test_parsers.py @@ -6,8 +6,6 @@ from datetime import datetime, timedelta from typing import Any -import pytest - from trendflow._parsers import ( _is_missing_value, _split_bracketed_ints, @@ -26,10 +24,7 @@ from trendflow.models import ( InterestByRegionResult, InterestOverTimeResult, - RegionalInterestRow, - RelatedQuery, RelatedResult, - TrendingItem, TrendingResult, ) diff --git a/tests/test_session.py b/tests/test_session.py index af33007..02776d1 100644 --- a/tests/test_session.py +++ b/tests/test_session.py @@ -2,7 +2,7 @@ from __future__ import annotations -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest