diff --git a/app/matcher.py b/app/matcher.py index 0db3973..beaf64a 100644 --- a/app/matcher.py +++ b/app/matcher.py @@ -5,7 +5,13 @@ from dataclasses import dataclass from pathlib import Path from app.aliases import ProductAliasIndex, ProductTypeMatch -from app.attributes import BrandIndex, Compatibility, compare_attributes, parse_attributes +from app.attributes import ( + BrandIndex, + Compatibility, + ParsedAttributes, + compare_attributes, + parse_attributes, +) from app.catalog import CatalogItem, build_catalog, load_raw_catalog from app.intents import detect_service_intent from app.models import Candidate, MatchResult @@ -31,8 +37,11 @@ class CatalogMatcher: type_index: dict[str, list[int]] = defaultdict(list) model_index: dict[str, set[int]] = defaultdict(set) + brand_item_index: dict[str, list[int]] = defaultdict(list) for index, item in enumerate(self.items): type_index[item.product_type].append(index) + if item.attributes.brand is not None: + brand_item_index[item.attributes.brand].append(index) for model_code in item.attributes.model_codes: model_index[model_code].add(index) self.type_index = { @@ -41,6 +50,9 @@ class CatalogMatcher: self.model_index = { model: frozenset(indexes) for model, indexes in model_index.items() } + self.brand_item_index = { + brand: tuple(indexes) for brand, indexes in brand_item_index.items() + } def match(self, message: str) -> MatchResult: if not message or not message.strip(): @@ -67,6 +79,9 @@ class CatalogMatcher: token_start=0, ) if product_match is None: + brand_result = self._match_brand_only(message, provisional_attributes) + if brand_result is not None: + return brand_result # Product evidence always has priority. Service intent detection # is only a diagnostic fallback; both unknown and service # messages share the required not_found response. @@ -132,6 +147,66 @@ class CatalogMatcher: candidates=response_candidates, ) + def _match_brand_only( + self, message: str, query_attributes: ParsedAttributes + ) -> MatchResult | None: + if ( + query_attributes.brand is None + or not query_attributes.brand_required + or query_attributes.model_codes + ): + return None + + candidate_indices = self.brand_item_index.get(query_attributes.brand, ()) + brand_match = self.brand_index.detect(message) + if not candidate_indices or brand_match is None: + return None + + compatible = { + index: Compatibility( + compatible=True, + matched_fields=1, + total_fields=1, + ) + for index in candidate_indices + } + lexical_scores = self.lexical_index.rank(message, "", candidate_indices) + ranked = self._apply_preferences( + lexical_scores, + compatible, + cheap=query_attributes.cheap_preference, + ) + if not ranked: + return None + if not query_attributes.cheap_preference: + # With no product type every item has the same brand evidence. + # Document length must not arbitrarily favor short catalog names. + ranked.sort(key=lambda candidate: candidate.item.sku) + + status = "matched" if len(ranked) == 1 else "ambiguous" + brand_signal = ProductTypeMatch( + canonical="", + score=brand_match.score, + exact=brand_match.score == 1.0, + alias=brand_match.alias, + token_start=brand_match.token_start, + ) + candidates = [ + Candidate( + sku=candidate.item.sku, + confidence=self._confidence( + candidate, + product_match=brand_signal, + candidate_count=len(ranked), + specificity=query_attributes.specificity, + status=status, + rank=rank, + ), + ) + for rank, candidate in enumerate(ranked[:3]) + ] + return MatchResult(message=message, status=status, candidates=candidates) + def _model_candidates(self, model_codes: frozenset[str]) -> frozenset[int]: if not model_codes: return frozenset() @@ -167,7 +242,7 @@ class CatalogMatcher: if price_range == 0 else float((maximum_price - item.price) / price_range) ) - ranking_score = lexical.score * 0.68 + cheap_score * 0.32 + ranking_score = cheap_score else: ranking_score = lexical.score ranked.append( @@ -179,12 +254,21 @@ class CatalogMatcher: ) ) - ranked.sort( - key=lambda candidate: ( - -candidate.ranking_score, - candidate.item.price if cheap else candidate.item.sku, + if cheap: + ranked.sort( + key=lambda candidate: ( + candidate.item.price, + -candidate.lexical.score, + candidate.item.sku, + ) + ) + else: + ranked.sort( + key=lambda candidate: ( + -candidate.ranking_score, + candidate.item.sku, + ) ) - ) return ranked @staticmethod diff --git a/tests/test_matcher.py b/tests/test_matcher.py index b85c819..cf1fa37 100644 --- a/tests/test_matcher.py +++ b/tests/test_matcher.py @@ -134,12 +134,49 @@ def test_makita_transliterations_share_the_same_brand_alias( assert result.status == "ambiguous" assert [candidate.sku for candidate in result.candidates] == [ - "INS-0010", "INS-0017", + "INS-0010", "INS-0003", ] +def test_cheap_preference_orders_candidates_by_ascending_price( + matcher: CatalogMatcher, +) -> None: + result = matcher.match("шуруповерт как у макиты, только дешевле") + prices_by_sku = {item.sku: item.price for item in matcher.items} + candidate_prices = [ + prices_by_sku[candidate.sku] for candidate in result.candidates + ] + + assert candidate_prices == sorted(candidate_prices) + + +@pytest.mark.parametrize( + "brand", ["ToolKraft", "toolkraft", "тулкрафт", "тул крафт"] +) +def test_brand_only_toolkraft_query_returns_catalog_items( + matcher: CatalogMatcher, brand: str +) -> None: + result = matcher.match(brand) + + assert result.status == "ambiguous" + assert [candidate.sku for candidate in result.candidates] == [ + "INS-0022", + "INS-0023", + "INS-0024", + ] + + +def test_unknown_toolkraft_model_does_not_fall_back_to_brand_results( + matcher: CatalogMatcher, +) -> None: + result = matcher.match("ToolKraft TK-999") + + assert result.status == "not_found" + assert result.candidates == [] + + @pytest.mark.parametrize("query", ["шурпуоверт на 12в", "шураыввавёрт на 12в"]) def test_product_type_survives_typos(matcher: CatalogMatcher, query: str) -> None: result = matcher.match(query) diff --git a/tests/test_message_corpus.py b/tests/test_message_corpus.py index b664a4d..48e0ef7 100644 --- a/tests/test_message_corpus.py +++ b/tests/test_message_corpus.py @@ -57,7 +57,7 @@ EXPECTED_RESULTS: dict[str, ExpectedResult] = { "лента фум 12 мм": ("matched", ("RAS-0060",)), "перчатки нитриловые есть?": ("matched", ("PER-0002",)), "проф труба 20х20 стенка полтора": ("matched", ("TRB-0002",)), - "шурик на 12в недорогой": ("ambiguous", ("INS-0010", "INS-0017", "INS-0003")), + "шурик на 12в недорогой": ("ambiguous", ("INS-0017", "INS-0010", "INS-0003")), "болгарка на 230 какая есть": ("ambiguous", ("INS-0014", "INS-0028", "INS-0007")), "гкл 9.5 сколько лист": ("matched", ("GKL-0001",)), "хомуты пластиковые 4.8х400": ("matched", ("RAS-0014",)), @@ -82,7 +82,7 @@ EXPECTED_RESULTS: dict[str, ExpectedResult] = { "сверло нужно": ("ambiguous", ("BIT-0060", "BIT-0061", "BIT-0062")), "какие есть диски": ("ambiguous", ("DSK-0011", "DSK-0012", "DSK-0013")), "перфоратор посоветуйте": ("ambiguous", ("INS-0012", "INS-0026", "INS-0005")), - "шуруповерт как у макиты, только дешевле": ("ambiguous", ("INS-0010", "INS-0017", "INS-0003")), + "шуруповерт как у макиты, только дешевле": ("ambiguous", ("INS-0017", "INS-0010", "INS-0003")), "здравствуйте, вы до скольки работаете?": ("not_found", ()), "можно оплатить картой при получении?": ("not_found", ()), "где находится ваш магазин": ("not_found", ()), @@ -120,7 +120,7 @@ EXPECTED_RESULTS: dict[str, ExpectedResult] = { "фанера 18 мм": ("not_found", ()), "дрель prowerk pw-750": ("matched", ("INS-0008",)), "аккумуляторная дрель 18в": ("not_found", ()), - "шуруповёрт 12 вольт самый дешёвый": ("ambiguous", ("INS-0010", "INS-0017", "INS-0003")), + "шуруповёрт 12 вольт самый дешёвый": ("ambiguous", ("INS-0017", "INS-0010", "INS-0003")), "перфоратор sds plus 800 вт": ("not_found", ()), "болгарка 125": ("ambiguous", ("INS-0013", "INS-0027", "INS-0006")), "ушм 230 мм 2200 вт": ("not_found", ()),