From 169f0ee72f8577db4c34e6711585c5e999b8b12e Mon Sep 17 00:00:00 2001 From: SAM LINUS Date: Thu, 6 Nov 2025 12:06:07 +0530 Subject: [PATCH] Included ai modules --- nance/ai/bert_helper.py | 25 +++ nance/ai/db_helper.py | 67 ++++++ nance/ai/db_query_helper.py | 148 +++++++++++++ nance/ai/flow.py | 413 ++++++++++++++++++++++++++++++++++++ nance/ai/hf.py | 97 +++++++++ nance/ai/llm.py | 33 +++ nance/ai/rule_extract.py | 255 ++++++++++++++++++++++ nance/ai/tavily_helper.py | 41 ++++ 8 files changed, 1079 insertions(+) create mode 100644 nance/ai/bert_helper.py create mode 100644 nance/ai/db_helper.py create mode 100644 nance/ai/db_query_helper.py create mode 100644 nance/ai/flow.py create mode 100644 nance/ai/hf.py create mode 100644 nance/ai/llm.py create mode 100644 nance/ai/rule_extract.py create mode 100644 nance/ai/tavily_helper.py diff --git a/nance/ai/bert_helper.py b/nance/ai/bert_helper.py new file mode 100644 index 00000000..0513e8e6 --- /dev/null +++ b/nance/ai/bert_helper.py @@ -0,0 +1,25 @@ +from hf import load_private_bert_model, set_pipeline, predict, group_fragments_by_gap, combine_words +from typing import List + +def run_bert(sms) -> List[str]: + try: + print("Running BERT...") + # Load the model and tokenizer + tokenizer, model = load_private_bert_model() + pipeline = set_pipeline(model, tokenizer) + response = predict(pipeline, sms) + # Group the fragments by gap + groups = group_fragments_by_gap(response) + if not groups: + return [] + # Combine the words + combined_words = combine_words(groups) + # Return the combined words + return combined_words + except Exception as e: + raise e + + + + + diff --git a/nance/ai/db_helper.py b/nance/ai/db_helper.py new file mode 100644 index 00000000..b8c0eb17 --- /dev/null +++ b/nance/ai/db_helper.py @@ -0,0 +1,67 @@ +import psycopg2 +from psycopg2.extras import DictCursor + +class PostgresDB: + def __init__(self, dbname="Nance", user="postgres", password="2003", host="localhost", port="5432"): + self.dbname = dbname + self.user = user + self.password = password + self.host = host + self.port = port + self.conn = None + self.cursor = None + self.schema = "nance_init" + + def connect(self): + try: + if self.conn is None or self.conn.closed: + self.conn = psycopg2.connect( + dbname=self.dbname, + user=self.user, + password=self.password, + host=self.host, + port=self.port + ) + print("Connection established to the DB ✅") + self.cursor = self.conn.cursor(cursor_factory=DictCursor) + print("Cursor created ✅") + return self.conn, self.cursor + except Exception as e: + raise e + + def execute_query(self, query, params=None): + try: + # Ensure DB connection + if self.conn is None or self.conn.closed: + self.connect() + + # Execute query + if params: + print(f"Executing query with params: {params}") + self.cursor.execute(query, params) + else: + print("Executing query without params") + self.cursor.execute(query) + + # Check if query produced rows (SELECT or RETURNING) + if self.cursor.description: + result = self.cursor.fetchall() + print(f"Query returned {len(result)} rows") + return result + + # For INSERT/UPDATE/DELETE without returning rows + self.conn.commit() + return None + except Exception as e: + print(f"Database error: {e}") + print(f"Query: {query}") + print(f"Params: {params}") + raise e + + def close(self): + try: + if self.conn and not self.conn.closed: + self.conn.close() + print("Connection closed to the DB ✅") + except Exception as e: + raise e \ No newline at end of file diff --git a/nance/ai/db_query_helper.py b/nance/ai/db_query_helper.py new file mode 100644 index 00000000..5ff98956 --- /dev/null +++ b/nance/ai/db_query_helper.py @@ -0,0 +1,148 @@ +from db_helper import PostgresDB +from typing import Tuple + + + +class DBQueryHelper: + def __init__(self): + #Initialize the database connection + self.db = PostgresDB() + self.db.connect() + # Initialize queries with proper schema formatting + self.initialize_queries() + self.transaction_type_map = { + "debit": "expense", + "credit": "income" + } + + def initialize_queries(self): + self.main_category_query = f""" + SELECT mc.main_category + FROM {self.db.schema}.merchants mer + JOIN {self.db.schema}.main_categories mc + ON mc.id = mer.main_category_id + WHERE mer.id = %s; + """ + self.sub_category_query = f""" + SELECT sc.sub_category + FROM {self.db.schema}.merchants mer + JOIN {self.db.schema}.sub_categories sc + ON sc.id = mer.sub_category_id + WHERE mer.id = %s; + """ + # Since transaction_type is from rules, the string will match exactly. + self.get_all_main_categories_query = f""" + SELECT DISTINCT id, main_category FROM {self.db.schema}.main_categories WHERE type = %s + """ + self.get_sub_category_by_main_category_query = f""" + SELECT sc.id, sc.sub_category + FROM {self.db.schema}.category_map cm + JOIN {self.db.schema}.sub_categories sc + ON cm.sub_category_id = sc.id + WHERE cm.main_category_id = %s; + """ + self.get_categories_by_merchant_id_query = f""" + SELECT mc.main_category + FROM {self.db.schema}.merchants mer + JOIN {self.db.schema}.main_categories mc + ON mc.id = mer.main_category_id + WHERE mer.id = %s; + """ + self.alike_merchant_names_query = f""" + SELECT + id, + merchant_name, + similarity(LOWER(merchant_name), LOWER(%s)) AS score + FROM {self.db.schema}.merchants + WHERE LOWER(merchant_name) %% LOWER(%s) + ORDER BY score DESC + LIMIT 10; + """ + #Close the database connection + def close_db_connection(self): + self.db.close() + + def get_sms_by_id(self, sms_id): + try: + # TODO: App inserts sms and app calls the pipeline. + # TODO: SMS is directly passed to the pipeline. + query = f"SELECT * FROM {self.db.schema}.sms WHERE id = %s" + self.db.execute_query(query, (sms_id,)) + return self.db.cursor.fetchone() + except Exception as e: + raise e + + + def get_all_main_categories(self, transaction_type: str) -> str: + try: + result = self.db.execute_query(self.get_all_main_categories_query, (self.transaction_type_map[transaction_type],)) + print(f"Main categories query result: {result}") + return result + except Exception as e: + raise e + + def get_sub_category_by_main_category(self, main_category: str) -> list: + try: + result = self.db.execute_query(self.get_sub_category_by_main_category_query, (main_category,)) + return result + except Exception as e: + raise e + + def get_categories_by_merchant_id(self, merchant_id: int) -> list: + try: + main_category = self.db.execute_query(self.main_category_query, (merchant_id,)) + if main_category: + # Get the sub category + sub_category = self.db.execute_query(self.sub_category_query, (merchant_id,)) + if sub_category: + # Both are present + return (main_category[0]["main_category"], sub_category[0]["sub_category"]) + else: + # Only main category is present + return (main_category[0]["main_category"], "None") + else: + # No category is present + return ("None", "None") + except Exception as e: + raise e + + def test_category_retrieval(self): + try: + query = f""" + SELECT sc.sub_category + FROM {self.db.schema}.merchants mer + JOIN {self.db.schema}.sub_categories sc + ON sc.id = mer.sub_category_id + WHERE mer.id = 1; + """ + result = self.db.execute_query(query) + print("None" if not result else result[0]["sub_category"]) + except Exception as e: + raise e + + def convert_None_to_empty_string(self, value: str) -> str: + return None if value == "" else value + + # Get the alike merchant names from the database. + def get_alike_merchant_names(self, merchant_name: str) -> str: + try: + # Ensure connection is established + self.db.connect() + print(f"Searching for merchants similar to: {merchant_name}") + result = self.db.execute_query(self.alike_merchant_names_query, (merchant_name,merchant_name)) + print(f"Result: {result}") + return result or [] + except Exception as e: + print(f"Error in get_alike_merchant_names: {e}") + return [] + + def get_sms_by_phone_number(self, phone_number: str) -> list: + try: + query = f"SELECT * FROM {self.db.schema}.sms WHERE phone_number = %s" + self.db.execute_query(query, (phone_number,)) + return self.db.cursor.fetchall() + except Exception as e: + raise e + +helper = DBQueryHelper() +helper.test_category_retrieval() \ No newline at end of file diff --git a/nance/ai/flow.py b/nance/ai/flow.py new file mode 100644 index 00000000..c4493a0f --- /dev/null +++ b/nance/ai/flow.py @@ -0,0 +1,413 @@ +from langchain_core.runnables import RunnableLambda +from prompt_templates import ( + NORMALIZE_PROMPT, + MAP_PROMPT, + SCRAPE_SUMMARIZE_PROMPT, + MAIN_CATEGORIZATION_PROMPT, + SUB_CATEGORY_PROMPT, + KNOWN_SUMMARIZE_PROMPT +) +from llm import llm, safe_extract_json +from bert_helper import run_bert +from db_query_helper import DBQueryHelper +from tavily_helper import TavilyHelper +from rule_extract import SMSExtractor +from typing import Literal, List, Optional, Union, Dict, Any +from dataclasses import dataclass, field, asdict +import warnings + +warnings.filterwarnings("ignore") + + +@dataclass +class MerchantData: + """Data class for merchant information and transaction details.""" + is_transaction: bool = False + transaction_type: str = "unknown" + amount: Optional[float] = None + payment_mode: str = "OTHERS" + transaction_date: Optional[str] = None + transaction_time: Optional[str] = None + candidates_list: List[str] = field(default_factory=list) + normalized_account_name: Optional[str] = None + confidence_score: float = 0.0 + main_category: str = "" + sub_category: str = "" + main_category_id: Optional[int] = None + sub_category_id: Optional[int] = None + context: str = "" + web_search_context: str = "" + matched_merchants: List[Dict[str, Any]] = field(default_factory=list) + merchant_name: str = "" + is_matched: bool = False + merchant_id: Optional[int] = None + content_type: str = "" + is_business: bool = False + main_categories: List[Dict[str, Any]] = field(default_factory=list) + sub_categories: List[Dict[str, Any]] = field(default_factory=list) + + def to_dict(self) -> Dict[str, Any]: + """Convert dataclass to dictionary.""" + # First convert any DictRow objects to regular dictionaries + self._convert_dict_rows() + return asdict(self) + + def _convert_dict_rows(self) -> None: + """Convert any DictRow objects to regular dictionaries to make them serializable.""" + # Check if matched_merchants contains DictRow objects + if hasattr(self, 'matched_merchants') and self.matched_merchants: + try: + # Convert DictRow objects to regular dictionaries + converted_merchants = [] + for merchant in self.matched_merchants: + if hasattr(merchant, 'keys'): # Check if it's a DictRow-like object + converted_merchants.append(dict(merchant)) + else: + converted_merchants.append(merchant) + self.matched_merchants = converted_merchants + except Exception as e: + print(f"Warning: Could not convert DictRow objects: {e}") + # Fallback: convert to empty list if conversion fails + self.matched_merchants = [] + + # Check other list fields that might contain DictRow objects + for field_name in ['main_categories', 'sub_categories']: + if hasattr(self, field_name) and getattr(self, field_name): + try: + field_value = getattr(self, field_name) + if isinstance(field_value, list): + converted_items = [] + for item in field_value: + if hasattr(item, 'keys'): # Check if it's a DictRow-like object + converted_items.append(dict(item)) + else: + converted_items.append(item) + setattr(self, field_name, converted_items) + except Exception as e: + print(f"Warning: Could not convert DictRow objects in {field_name}: {e}") + setattr(self, field_name, []) + + def update_from_dict(self, data: Dict[str, Any]) -> None: + """Update fields from a dictionary, only updating existing fields.""" + for key, value in data.items(): + if hasattr(self, key): + setattr(self, key, value) + + +class SMSMerchantCategorizer: + """Main class for categorizing SMS merchants and transactions.""" + + def __init__(self): + self.merchant_data = MerchantData() + self.db_query_helper = DBQueryHelper() + self.tavily_helper = TavilyHelper() + self.sms_extractor = SMSExtractor() + + def get_sms_details(self, sms: str) -> None: + """Extract SMS details and update merchant data.""" + try: + sms_details = self.sms_extractor.extract_details(sms) + if sms_details: + self.merchant_data.update_from_dict(sms_details) + except Exception as e: + print(f"Error extracting SMS details: {e}") + + def categorization_flow(self, content: str, content_type: str) -> Dict[str, Any]: + """Main categorization flow for processing content.""" + try: + # Reset merchant data for new content + self.merchant_data = MerchantData() + self.merchant_data.content_type = content_type + + # Extract details from SMS if applicable + if content_type == "SMS": + self.get_sms_details(content) + + # Early return if SMS is not a transaction + if not self.merchant_data.is_transaction: + print("SMS is not a transaction ⚠️") + return self.merchant_data.to_dict() + + # Extract candidates using BERT + candidates = run_bert(content) + self.merchant_data.candidates_list = candidates if candidates else [] + + if not self.merchant_data.candidates_list: + print("No candidates found ⚠️") + return self.merchant_data.to_dict() + + print(f"Candidates list: {self.merchant_data.candidates_list}") + + # Normalize account name + response_json = self._run_normalize_chain() + if response_json: + self.merchant_data.update_from_dict(response_json) + + print(f"Merchant data: {self.merchant_data.to_dict()}") + + # Check if business account + self._convert_bool_field("is_business") + if not self.merchant_data.is_business: + print(f"Normalized account name: {self.merchant_data.normalized_account_name}") + print("Account name is not a merchant ⚠️") + return self.merchant_data.to_dict() + + if not self.merchant_data.normalized_account_name: + print("No account name found ⚠️") + return self.merchant_data.to_dict() + + # Get matched merchants from database + self._process_matched_merchants() + + # Process merchant matching + self._convert_bool_field("is_matched") + if self.merchant_data.is_matched: + self._process_merchant_categories() + return self.merchant_data.to_dict() + + # Handle unknown merchants + if self.merchant_data.confidence_score < 75: + self._process_unknown_merchant() + + else: + response_json = self._run_known_summarize_chain() + self.merchant_data.context = response_json.get("summary", "") + + # Get main categories based on transaction type + self._get_main_categories() + + if not self.merchant_data.main_categories: + print("No main categories found ⚠️") + return self.merchant_data.to_dict() + + print("Main categories found ✅") + + # Categorize to main category + response_json = self._run_main_categorization_chain() + self.merchant_data.main_category = response_json.get("main_category", "") + self.merchant_data.main_category_id = response_json.get("main_category_id") + + if not self.merchant_data.main_category or self.merchant_data.main_category == "None": + print("No main category found ⚠️") + return self.merchant_data.to_dict() + + print(f"Main category: {self.merchant_data.main_category}") + print(f"Main category ID: {self.merchant_data.main_category_id}") + + # Process sub-categories + self._process_sub_categories() + + return self.merchant_data.to_dict() + + except Exception as e: + print(f"Error running flow: {e}") + import traceback + print(f"Traceback: {traceback.format_exc()}") + return {} + + def _process_matched_merchants(self) -> None: + """Process matched merchants from database.""" + try: + matched_merchants = self.db_query_helper.get_alike_merchant_names( + self.merchant_data.normalized_account_name + ) + print("Alike merchants:", matched_merchants) + + if matched_merchants: + print("Matched merchants found ✅") + self.merchant_data.matched_merchants = matched_merchants + response_json = self._run_map_chain() + print("Mapped merchant to DB ✅") + + if response_json: + self.merchant_data.update_from_dict(response_json) + else: + print("No matched merchants found ⚠️") + self.merchant_data.matched_merchants = [] + self.merchant_data.merchant_name = self.merchant_data.normalized_account_name + self.merchant_data.is_matched = False + self.merchant_data.merchant_id = None + + except Exception as e: + print(f"Database query error: {e}") + self.merchant_data.matched_merchants = [] + self.merchant_data.is_matched = False + self.merchant_data.merchant_id = None + + def _process_merchant_categories(self) -> None: + """Process merchant categories for matched merchants.""" + try: + merchant_category = self.db_query_helper.get_categories_by_merchant_id( + self.merchant_data.merchant_id + ) + if merchant_category and len(merchant_category) >= 2: + self.merchant_data.main_category = merchant_category[0] + self.merchant_data.sub_category = merchant_category[1] + else: + print("Invalid merchant category data") + self.merchant_data.is_matched = False + except Exception as e: + print(f"Error getting merchant categories: {e}") + self.merchant_data.is_matched = False + + def _process_unknown_merchant(self) -> None: + """Process unknown merchants by getting context.""" + self.merchant_data.context = "" + # Get merchant context by summarizing web searches + merchant_context = self.tavily_helper.handle_web_search( + self.merchant_data.merchant_name + ) + self.merchant_data.web_search_context = merchant_context + + response_json = self._run_web_search_summarize_chain() + self.merchant_data.context = response_json.get("summary", "") + # TODO: Insert merchant name and description as new row in table + # TODO: Insert merchant's category and sub-category to DB + + def _get_main_categories(self) -> None: + """Get main categories based on transaction type.""" + if self.merchant_data.transaction_type.lower() != "unknown": + self.merchant_data.main_categories = self.db_query_helper.get_all_main_categories( + self.merchant_data.transaction_type.lower() + ) + else: + self.merchant_data.main_categories = [] + + def _process_sub_categories(self) -> None: + """Process sub-categories for the identified main category.""" + sub_categories = self.db_query_helper.get_sub_category_by_main_category( + self.merchant_data.main_category_id + ) + self.merchant_data.sub_categories = sub_categories if sub_categories else [] + + response_json = self._run_sub_category_chain() + self.merchant_data.sub_category = response_json.get("sub_category", "") + self.merchant_data.sub_category_id = response_json.get("sub_category_id") + + def _run_main_categorization_chain(self) -> Dict[str, Any]: + """Run the main categorization chain.""" + try: + main_categorization_chain = ( + MAIN_CATEGORIZATION_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + response_json = main_categorization_chain.invoke({ + "merchant_name": self.merchant_data.merchant_name, + "merchant_context": ( + self.merchant_data.context + if self.merchant_data.confidence_score < 50 + else "" + ), + "categories": self.merchant_data.main_categories, + "transaction_type": self.merchant_data.transaction_type + }) + return response_json or {} + except Exception as e: + print(f"Error in main categorization chain: {e}") + return {} + + def _run_sub_category_chain(self) -> Dict[str, Any]: + """Run the sub-category chain.""" + try: + sub_category_chain = ( + SUB_CATEGORY_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + response_json = sub_category_chain.invoke({ + "merchant_name": self.merchant_data.merchant_name, + "merchant_context": self.merchant_data.context, + "main_category": self.merchant_data.main_category, + "sub_categories": self.merchant_data.sub_categories, + "transaction_type": self.merchant_data.transaction_type + }) + return response_json or {} + except Exception as e: + print(f"Error in sub category chain: {e}") + return {} + + def _run_web_search_summarize_chain(self) -> Dict[str, Any]: + """Run the web search summarize chain.""" + try: + summarize_chain = ( + SCRAPE_SUMMARIZE_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + response_json = summarize_chain.invoke({ + "merchant_name": self.merchant_data.merchant_name, + "merchant_context": self.merchant_data.web_search_context + }) + return response_json or {} + except Exception as e: + print(f"Error in web search summarize chain: {e}") + return {} + + def _run_known_summarize_chain(self) -> Dict[str, Any]: + """Run the known summarize chain.""" + try: + summarize_chain = ( + KNOWN_SUMMARIZE_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + prompt_str = KNOWN_SUMMARIZE_PROMPT.format_prompt( + merchant_name=self.merchant_data.merchant_name + ).to_string() + print("Final Prompt:\n", prompt_str) + response_json = summarize_chain.invoke({ + "merchant_name": self.merchant_data.merchant_name, + }) + return response_json or {} + except Exception as e: + print(f"Error in known summarize chain: {e}") + return {} + + def _convert_bool_field(self, key: str) -> None: + """Convert string boolean fields to actual boolean values.""" + if hasattr(self.merchant_data, key): + value = getattr(self.merchant_data, key) + if isinstance(value, str): + if value.lower() == "true": + setattr(self.merchant_data, key, True) + elif value.lower() == "false": + setattr(self.merchant_data, key, False) + + def _run_map_chain(self) -> Dict[str, Any]: + """Run the map chain.""" + try: + map_chain = ( + MAP_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + response_json = map_chain.invoke({ + "merchant_name": self.merchant_data.normalized_account_name, + "matched_merchants": self.merchant_data.matched_merchants + }) + return response_json or {} + except Exception as e: + print(f"Error in map chain: {e}") + return {} + + def _run_normalize_chain(self) -> Dict[str, Any]: + """Run the normalize chain.""" + try: + normalize_chain = ( + NORMALIZE_PROMPT | + llm | + RunnableLambda(lambda x: safe_extract_json(x.content)) + ) + response_json = normalize_chain.invoke({ + "candidates_list": self.merchant_data.candidates_list + }) + return response_json or {} + except Exception as e: + print(f"Error in normalize chain: {e}") + return {} + + +if __name__ == "__main__": + categorizer = SMSMerchantCategorizer() + # categorizer.categorization_flow("I paid 1000 to Amazon for groceries") \ No newline at end of file diff --git a/nance/ai/hf.py b/nance/ai/hf.py new file mode 100644 index 00000000..252784d2 --- /dev/null +++ b/nance/ai/hf.py @@ -0,0 +1,97 @@ +import os, torch, warnings +from dotenv import load_dotenv +from huggingface_hub import login +from transformers import AutoTokenizer, AutoModelForTokenClassification, pipeline +from dataset import read_data +from tqdm import tqdm +warnings.filterwarnings("ignore") + +MODEL_NAME = "Samlinus/NANCE_BERT_NER_AA_MODEL" +LOCAL_MODEL_DIR = "./local_nance_bert_ner" + +def load_private_bert_model(model_name: str = MODEL_NAME, local_dir: str = LOCAL_MODEL_DIR) -> tuple: + # Load token from .env + load_dotenv() + token = os.getenv("HUGGINGFACE_TOKEN") + if not token: + raise ValueError("HUGGINGFACE_TOKEN not found in .env file ❌") + + # If local model exists, load from local directory + if os.path.exists(local_dir) and os.path.isdir(local_dir): + print("Loading model and tokenizer from local directory...") + tokenizer = AutoTokenizer.from_pretrained(local_dir) + model = AutoModelForTokenClassification.from_pretrained(local_dir) + print(f"Model and tokenizer loaded from {local_dir} ✅") + else: + print("Downloading model and tokenizer from Hugging Face...") + # Login to Hugging Face + login(token=token) + # Download and save model/tokenizer locally + tokenizer = AutoTokenizer.from_pretrained(model_name, use_auth_token=True) + model = AutoModelForTokenClassification.from_pretrained(model_name, use_auth_token=True) + # Save model and tokenizer locally + tokenizer.save_pretrained(local_dir) + model.save_pretrained(local_dir) + print(f"Model and tokenizer saved to {local_dir} ✅") + return tokenizer, model + +def set_pipeline(model, tokenizer) -> pipeline: + return pipeline( + "ner", + model=model, + tokenizer=tokenizer, + device=0 if torch.cuda.is_available() else -1, + aggregation_strategy="simple" +) + +def predict(pipeline, sms) -> list: + return pipeline(sms) + + +def group_fragments_by_gap(bert_response, max_gap=5) -> list: + if len(bert_response) == 0: + return [] + bert_response = sorted(bert_response, key=lambda x: x['start']) + groups = [] + current_group = [bert_response[0]] + for i in range(1, len(bert_response)): + prev_end = current_group[-1]['end'] + curr_start = bert_response[i]['start'] + if curr_start - prev_end <= max_gap: + current_group.append(bert_response[i]) + else: + groups.append(current_group) + current_group = [bert_response[i]] + groups.append(current_group) + return groups + +def combine_words(groups) -> list: + return [''.join([frag['word'] for frag in group]).strip() for group in groups] + + +def run_pipeline(): + try: + # Reading `test_sms_aa.csv` + df = read_data() + # Loading private model (now with local caching) + tokenizer, model = load_private_bert_model(MODEL_NAME, LOCAL_MODEL_DIR) + # Setting pipeline + pipeline = set_pipeline(model, tokenizer) + # Predicting NER + df["groups"] = df["Input"].apply(lambda x: predict(pipeline, x)) + # Grouping fragments by gap + df["groups"] = df["groups"].apply(lambda x: group_fragments_by_gap(x)) + # Combining words + df["combined_words"] = df["groups"].apply(lambda x: combine_words(x)) + # Saving to `test_sms_aa_groups.csv` + df.to_csv("test_sms_aa_groups.csv", index=False) + print(f"Data saved to `test_sms_aa_groups.csv` ✅") + except Exception as e: + print(f"Error running pipeline: {e.with_traceback()}") + print(f"Pipeline failed ❌ ") + + +# Usage example: +if __name__ == "__main__": + run_pipeline() + diff --git a/nance/ai/llm.py b/nance/ai/llm.py new file mode 100644 index 00000000..b61f118a --- /dev/null +++ b/nance/ai/llm.py @@ -0,0 +1,33 @@ +import json +import re +from langchain_core.runnables import RunnableLambda, RunnableBranch +from langchain_groq import ChatGroq +from dotenv import load_dotenv +load_dotenv() + +def safe_extract_json(text): + """ + Safely extract a JSON object from a string, handling code blocks and partial JSON. + """ + try: + return json.loads(text) + except json.JSONDecodeError: + # Extract JSON block between ```json ... ``` + match = re.search(r"```json\s*(\{.*?\})\s*```", text, re.DOTALL) + if match: + try: + return json.loads(match.group(1)) + except json.JSONDecodeError: + pass + # Try any JSON-looking object in the string + match = re.search(r"(\{.*\})", text, re.DOTALL) + if match: + try: + return json.loads(match.group(1)) + except json.JSONDecodeError: + pass + raise ValueError("Failed to parse LLM response as JSON.") + + +# Expose a default LLM as a Runnable for chaining +llm = ChatGroq(model="qwen/qwen3-32b", temperature=0.2, max_retries=1) \ No newline at end of file diff --git a/nance/ai/rule_extract.py b/nance/ai/rule_extract.py new file mode 100644 index 00000000..a8c566b1 --- /dev/null +++ b/nance/ai/rule_extract.py @@ -0,0 +1,255 @@ +import re +from typing import Dict, Optional, List, Tuple, Literal, Any, Union +from datetime import datetime + + +class SMSExtractor: + """ + A class for extracting transaction details from SMS messages. + + This class provides methods to identify and extract various transaction-related + information from SMS messages including transaction type, amount, and payment mode. + """ + + def __init__(self): + """Initialize the SMSExtractor with predefined keyword lists.""" + self.transaction_keywords = [ + "debited", "credited", "spent", "withdrawn", "purchase", + "deposited", "sent", "transferred", "received", "neft", + "cr", "txn", "transaction", "paid" + ] + + self.debit_keywords = [ + "debited", "spent", "withdrawn", "purchase", + "transferred", "paid", "sent" + ] + + self.credit_keywords = [ + "credited", "deposited", "received", "cr" + ] + + self.payment_mode_patterns = { + "UPI": ["upi", "paytm", "gpay", "bhim", "bhimupi", "phonepe"], + "ATM": ["atm"], + "BANK TRANSFER": ["neft", "rtgs", "imps"], + "CARD": ["debit card", "card ending"] + } + + + def is_transaction(self, sms: str) -> bool: + """ + Determine if an SMS contains transaction information. + + Args: + sms (str): The SMS message to analyze + + Returns: + bool: True if transaction keywords are found, False otherwise + + Raises: + ValueError: If SMS is not a string + """ + try: + return True if any(kw in sms.lower() for kw in self.transaction_keywords) else False + except Exception as e: + print(f"Error in is_transaction: {e}") + raise e + + def get_transaction_type(self, sms: str) -> str: + """ + Determine the type of transaction (debit/credit/unknown). + + Args: + sms (str): The SMS message to analyze + + Returns: + str: "debit", "credit", or "unknown" + """ + try: + sms_lower = sms.lower() + + # High-confidence override + if "neft" in sms_lower: + return "credit" + + has_debit = any(word in sms_lower for word in self.debit_keywords) + has_credit = any(word in sms_lower for word in self.credit_keywords) + + if has_debit: + return "debit" + if has_credit: + return "credit" + return "unknown" # if no keywords are found + except Exception as e: + print(f"Error in get_transaction_type: {e}") + raise e + + def extract_amount(self, sms: str) -> Optional[str]: + """ + Extract the transaction amount from the SMS. + + Args: + sms (str): The SMS message to analyze + + Returns: + Optional[str]: The extracted amount as string, or None if not found + """ + try: + # Extract transaction amount using regex + amount_match = re.search(r'(?:rs\.?|inr)[\s:.]*([\d,]+(?:\.\d{1,2})?)', sms.lower()) + if amount_match: + return float(amount_match.group(1).replace(",", "")) + return None + except Exception as e: + print(f"Error in extract_amount: {e}") + raise e + + def get_payment_mode(self, sms: str) -> Literal["UPI", "ATM", "BANK TRANSFER", "CARD", "OTHERS"]: + """ + Identify the payment mode used in the transaction. + + Args: + sms (str): The SMS message to analyze + + Returns: + Literal["UPI", "ATM", "BANK TRANSFER", "CARD", "OTHERS"]: The payment mode + """ + try: + sms_lower = sms.lower() + + for mode, patterns in self.payment_mode_patterns.items(): + if any(pattern in sms_lower for pattern in patterns): + return mode + + return "OTHERS" + except Exception as e: + print(f"Error in get_payment_mode: {e}") + raise e + + def extract_details(self, sms: str) -> Dict[str, Any]: + """ + Extract all transaction details from an SMS message. + + Args: + sms (str): The SMS message to analyze + + Returns: + Dict[str, Any]: Dictionary containing extracted details: + - is_transaction: True or False + - transaction_type: "debit", "credit", or "unknown" + - amount: Transaction amount or None + - payment_mode: Payment mode (UPI, ATM, BANK TRANSFER, CARD, OTHERS) + - transaction_date: Transaction date or None + - transaction_time: Transaction time or None + """ + try: + transaction_date, transaction_time = self.extract_transaction_date_and_time(sms) + result = { + "is_transaction": self.is_transaction(sms) + } + if self.is_transaction(sms): + result["transaction_type"] = self.get_transaction_type(sms) + result["amount"] = self.extract_amount(sms) + result["payment_mode"] = self.get_payment_mode(sms) + result["transaction_date"] = transaction_date + result["transaction_time"] = transaction_time + + print("Details extracted successfully ✅") + return result + except Exception as e: + print(f"Error in extracting details ❌: {e}") + raise e + + def extract_transaction_date_and_time(self, sms: str) -> Tuple[Optional[str], Optional[str]]: + """ + Extract the transaction date and time from the SMS. + + Args: + sms (str): The SMS message to analyze + + Returns: + Tuple[Optional[str], Optional[str]]: The extracted date and time + """ + sms = sms.lower() + + # --- DATE Patterns --- + date_patterns = [ + # Represents different date formats + (r"(\d{1,2})[/-](\d{1,2})[/-](\d{2,4})", ["%d-%m-%y", "%d-%m-%Y", "%d/%m/%y", "%d/%m/%Y"]), + (r"(\d{1,2})[/-]([a-z]{3})[/-](\d{2,4})", ["%d-%b-%y", "%d-%b-%Y", "%d/%b/%y", "%d/%b/%Y"]) + ] + + # --- TIME Pattern (with optional seconds and AM/PM) --- + time_regex = r"(\d{1,2}:\d{2}(?::\d{2})?) ?(am|pm)?" + + extracted_date = None + extracted_time = None + + for date_pattern, formats in date_patterns: + match = re.search(date_pattern, sms) + if match: + day, month, year = match.groups() + # Try all the formats + for fmt in formats: + try: + y = year + if len(y) == 2 and "%Y" in fmt: + y = str(datetime.now().year)[:2] + y + date_str = f"{day}-{month}-{y}" + dt = datetime.strptime(date_str, fmt) + extracted_date = dt.strftime("%Y-%m-%d") + + # Look for time near the matched date + # From the patterns observed if time occurs, it seems to be appearing closer to the date. + start = max(match.start() - 5, 0) + end = min(match.end() + 5, len(sms)) + nearby_text = sms[start:end] + + time_match = re.search(time_regex, nearby_text) + if time_match: + time_val, meridian = time_match.groups() + if meridian: + fmt = "%I:%M:%S %p" if time_val.count(":") == 2 else "%I:%M %p" + time_dt = datetime.strptime(f"{time_val} {meridian}", fmt) + else: + fmt = "%H:%M:%S" if time_val.count(":") == 2 else "%H:%M" + time_dt = datetime.strptime(time_val, fmt) + extracted_time = time_dt.strftime("%H:%M") + break + except: + continue + if extracted_date: + break + + return extracted_date, extracted_time + +# Backward compatibility functions +def is_transaction(sms: str) -> str: + """Backward compatibility function.""" + extractor = SMSExtractor() + return extractor.is_transaction(sms) + + +def transaction_type(sms: str) -> str: + """Backward compatibility function.""" + extractor = SMSExtractor() + return extractor.get_transaction_type(sms) + + +def amount(sms: str) -> str: + """Backward compatibility function.""" + extractor = SMSExtractor() + return extractor.extract_amount(sms) + + +def payment_mode(sms: str) -> str: + """Backward compatibility function.""" + extractor = SMSExtractor() + return extractor.get_payment_mode(sms) + + +def extract_details(sms: str) -> dict: + """Backward compatibility function.""" + extractor = SMSExtractor() + return extractor.extract_details(sms) + diff --git a/nance/ai/tavily_helper.py b/nance/ai/tavily_helper.py new file mode 100644 index 00000000..72faf7dd --- /dev/null +++ b/nance/ai/tavily_helper.py @@ -0,0 +1,41 @@ +from typing import Any + + +from langchain.tools import tool +from langchain_community.tools.tavily_search import TavilySearchResults + +class TavilyHelper: + def __init__(self): + self.search_tool = TavilySearchResults(k=1) + + def handle_web_search(self, merchant: str) -> str: + try: + query = f"About {merchant} company" + result = self.search_tool.run(query) + print("Tavily search result Done ✅") + top_k_contents = self.get_top_k_contents(result) + formatted_context = self.format_context(top_k_contents) + return formatted_context + except Exception as e: + raise e + + def trim_content(self, text: str, max_words: int = 150) -> str: + return " ".join(text.strip().split()[:max_words]) + + def format_context(self, top_results: list) -> str: + try: + formatted_results = "" + for i, result in enumerate(top_results): + title = result['title'].strip() + content = self.trim_content(result['content'], max_words=150) + formatted_results += f"{i+1}. Title: {title}\nContent: {content}\n\n" + return formatted_results + except Exception as e: + raise e + + def get_top_k_contents(self, results: list, k: int = 3) -> list: + try: + top_results = sorted(results, key=lambda x: x["score"], reverse=True)[:k] + return top_results + except Exception as e: + raise e \ No newline at end of file