File size: 6,124 Bytes
db4ba8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""
TradeFlow AI — HS Code RAG Service (Phase 3, Step 3.1)

PRD §12 — Retrieval-Augmented Generation for BTKI HS Code classification.

Architecture:
  1. text-embedding-3-small → embed product description
  2. ChromaDB → semantic search over 12,000+ HS codes
  3. Gemini Flash → re-rank top-10 candidates → return top-3 with reasoning
"""

from __future__ import annotations

import chromadb
import structlog
from chromadb.utils import embedding_functions
from langchain_core.messages import HumanMessage
from langchain_google_genai import ChatGoogleGenerativeAI

from ..config import settings

log = structlog.get_logger()

COLLECTION_NAME = "btki_hs_codes"

# ── Reranking prompt ──────────────────────────────────────────────────────────
RERANK_PROMPT = """
Anda adalah pakar bea cukai Indonesia yang ahli dalam Buku Tarif Kepabeanan Indonesia (BTKI).

Deskripsi produk: {description}

Kandidat kode HS berdasarkan pencarian semantik:
{candidates}

Tugas: Pilih 3 kode HS yang paling akurat. Berikan alasan singkat dalam Bahasa Indonesia.
Jawab dalam format JSON: [{{"hs_code": "XXXX.XX.XX", "description": "...", "reason": "...", "confidence": 0.XX}}]
"""


class HSRecommendService:
    """
    HS Code recommendation using RAG over BTKI vector store.
    """

    def __init__(self) -> None:
        # ChromaDB client
        self.chroma = chromadb.HttpClient(
            host=settings.CHROMADB_HOST,
            port=settings.CHROMADB_PORT,
        )

        # Google Gemini embedding function (free, text-embedding-004)
        self.embed_fn = embedding_functions.GoogleGenerativeAiEmbeddingFunction(
            api_key=settings.GEMINI_API_KEY,
            model_name=settings.EMBEDDING_MODEL,
        )

        # Gemini reranker
        primary_llm = ChatGoogleGenerativeAI(
            model=settings.GEMINI_MODEL_PRIMARY,
            temperature=0.1,
            api_key=settings.GEMINI_API_KEY,
        )
        fallback_llm = ChatGoogleGenerativeAI(
            model=settings.GEMINI_MODEL_FALLBACK,
            temperature=0.1,
            api_key=settings.GEMINI_API_KEY,
        )
        self.llm = primary_llm.with_fallbacks([fallback_llm])

    def _get_collection(self) -> chromadb.Collection:
        return self.chroma.get_or_create_collection(
            name=COLLECTION_NAME,
            embedding_function=self.embed_fn,
            metadata={"hnsw:space": "cosine"},
        )

    async def recommend(
        self,
        product_description: str,
        top_k: int = 10,
        return_count: int = 3,
    ) -> list[dict]:
        """
        Returns top `return_count` HS code recommendations.
        """
        log.info("HS Recommend request", description=product_description[:80])

        collection = self._get_collection()

        # Step 1 — Semantic search
        results = collection.query(
            query_texts=[product_description],
            n_results=top_k,
            include=["documents", "metadatas", "distances"],
        )

        candidates_raw = []
        if results["documents"] and results["documents"][0]:
            for doc, meta, dist in zip(
                results["documents"][0],
                results["metadatas"][0],
                results["distances"][0], strict=False,
            ):
                candidates_raw.append(
                    f"- HS {meta.get('hs_code', '?')}: {doc} "
                    f"(similarity {round(1 - dist, 3)})"
                )

        # Step 2 — LLM rerank
        candidates_text = "\n".join(candidates_raw) or "(tidak ada kandidat ditemukan)"
        prompt = RERANK_PROMPT.format(
            description=product_description,
            candidates=candidates_text,
        )

        try:
            import json
            response = await self.llm.ainvoke([HumanMessage(content=prompt)])
            # Parse JSON from response
            raw = response.content.strip()
            if raw.startswith("```"):
                raw = raw.split("```")[1].lstrip("json").strip()
            recommendations = json.loads(raw)[:return_count]
        except Exception as exc:
            log.error("LLM reranking failed", error=str(exc))
            # Fallback: return raw chromadb results
            recommendations = [
                {
                    "hs_code": r["metadatas"][0][i].get("hs_code", "0000.00.00") if r["metadatas"] and r["metadatas"][0] else "0000.00.00",
                    "description": r["documents"][0][i] if r["documents"] and r["documents"][0] else "",
                    "reason": "Pencarian semantik (reranking gagal)",
                    "confidence": round(1 - r["distances"][0][i], 3) if r["distances"] and r["distances"][0] else 0.0,
                }
                for i, r in enumerate([results] * min(return_count, top_k))
            ][:return_count]

        log.info("HS Recommend result", count=len(recommendations))
        return recommendations

    async def ingest_btki(self, records: list[dict]) -> int:
        """
        Ingest BTKI records into ChromaDB.
        Expected format: [{hs_code, description_id, description_en, duty_rate, ...}]
        """
        collection = self._get_collection()

        documents = [f"{r['hs_code']}: {r['description_id']} / {r['description_en']}" for r in records]
        metadatas = [
            {
                "hs_code": r["hs_code"],
                "duty_rate": str(r.get("duty_rate", 0)),
                "vat_rate": str(r.get("vat_rate", 0.11)),
            }
            for r in records
        ]
        ids = [r["hs_code"] for r in records]

        collection.upsert(documents=documents, metadatas=metadatas, ids=ids)
        log.info("BTKI ingested", count=len(records))
        return len(records)


# ── Singleton ────────────────────────────────────────────────────────────────
hs_recommend_service = HSRecommendService()