File size: 9,136 Bytes
1cd53ba bc2df26 1cd53ba 4ce08fd 87c0663 1cd53ba 87c0663 1cd53ba 4ce08fd 1cd53ba bc2df26 1cd53ba bc2df26 1cd53ba 64075d7 1cd53ba 64075d7 1cd53ba 4ce08fd 1cd53ba bc2df26 1cd53ba bc2df26 1cd53ba bc2df26 1cd53ba | 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 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 | import json
import logging
import re
from datetime import date
from pathlib import Path
from typing import Any
from app.ai.orchestrator import AIOrchestrator
from app.ai.router import LiteLLMOrchestration
from app.config import Settings
from app.database.supabase import SupabaseRepository
from app.models.domain import ExtractedTrip, WhatsAppInboundMessage
from app.services.embedding_service import JinaEmbeddingService
from app.services.trip_indexing import index_trip
from app.utils.departure import normalize_departure_bucket, _parse_date_value
from app.utils.time import now_in_timezone
logger = logging.getLogger(__name__)
_EXTRACTION_PROMPT_PATH = Path("prompts/group_trip_extraction.md")
class GroupMessageService:
def __init__(
self,
*,
repository: SupabaseRepository,
embeddings: JinaEmbeddingService,
ai: AIOrchestrator | LiteLLMOrchestration,
settings: Settings,
) -> None:
self.repository = repository
self.embeddings = embeddings
self.ai = ai
self.settings = settings
async def handle_group_message(
self,
inbound: WhatsAppInboundMessage,
) -> None:
if await self.repository.message_exists(inbound.message_id):
logger.info("Skipping duplicate group message %s", inbound.message_id)
return
extracted = await self._extract_trip_from_text(inbound.text)
if not extracted or not extracted.is_trip_ad:
logger.info("Group message %s is not a trip advertisement", inbound.message_id)
return
if not self._validate_extracted_trip(extracted):
logger.warning(
"Group message %s has incomplete trip data: %s",
inbound.message_id,
extracted,
)
return
departure_time = normalize_departure_bucket(extracted.departure_time)
if not departure_time:
logger.warning(
"Group message %s: could not normalize departure_time '%s'",
inbound.message_id,
extracted.departure_time,
)
return
departure_date = _parse_date_value(extracted.departure_date)
if not departure_date:
logger.warning(
"Group message %s: invalid departure_date '%s'",
inbound.message_id,
extracted.departure_date,
)
return
phone = self._normalize_phone(
extracted.driver_phone,
country_code=self.settings.default_country_code,
)
if not phone:
logger.warning(
"Group message %s: no valid phone number extracted",
inbound.message_id,
)
return
existing_customer = await self.repository.get_customer_by_phone_number(phone)
if existing_customer:
if existing_customer.get("registered"):
logger.info(
"Driver phone %s is registered (customer %s); discarding group message %s",
phone,
existing_customer["id"],
inbound.message_id,
)
return
driver = await self.repository.get_driver_by_phone_number(phone)
if driver:
existing_trip = await self.repository.get_driver_trip_by_datetime(
driver_id=str(driver["id"]),
departure_date=departure_date,
departure_time=departure_time,
)
if existing_trip:
logger.info(
"Driver %s already has trip at %s %s; discarding group message %s",
phone,
departure_date,
departure_time,
inbound.message_id,
)
return
trip = await self.repository.create_unregistered_driver_trip(
driver_id=str(driver["id"]) if driver else None,
phone_number=phone,
driver_name=extracted.driver_name,
car_type=extracted.car_type,
departure=extracted.departure,
destination=extracted.destination,
departure_date=departure_date,
departure_time=departure_time,
available_seats=extracted.available_seats,
total_seats=extracted.total_seats,
price=extracted.price or 0,
)
else:
trip = await self.repository.create_unregistered_driver_entities(
phone_number=phone,
driver_name=extracted.driver_name,
car_type=extracted.car_type,
departure=extracted.departure,
destination=extracted.destination,
departure_date=departure_date,
departure_time=departure_time,
available_seats=extracted.available_seats,
total_seats=extracted.total_seats,
price=extracted.price or 0,
)
await index_trip(
repository=self.repository,
embeddings=self.embeddings,
embedding_model=self.settings.jina_embedding_model,
trip=trip,
)
logger.info(
"Created trip %s from group message (unregistered driver %s)",
trip.get("id"),
phone,
)
async def _extract_trip_from_text(self, text: str) -> ExtractedTrip | None:
prompt_template = _EXTRACTION_PROMPT_PATH.read_text(encoding="utf-8")
dt = now_in_timezone(self.settings.app_timezone)
prompt = prompt_template.format(current_datetime=dt.isoformat())
try:
response = await self.ai.chat(
messages=[
{"role": "system", "content": prompt},
{"role": "user", "content": text},
],
tools=None,
tool_choice=None,
temperature=0.1,
)
except Exception as exc:
logger.warning("LLM extraction failed for group message: %s", exc)
return None
content = (response.content or "").strip()
if not content:
return None
return self._parse_extracted_json(content)
def _parse_extracted_json(self, content: str) -> ExtractedTrip | None:
if content.startswith("```"):
lines = content.split("\n")
content = "\n".join(lines[1:-1])
content = content.strip()
try:
data = json.loads(content)
except json.JSONDecodeError:
logger.warning("Failed to parse LLM extraction JSON: %s", content[:200])
return None
if not isinstance(data, dict):
return None
return ExtractedTrip(
is_trip_ad=bool(data.get("is_trip_ad", False)),
departure=data.get("departure"),
destination=data.get("destination"),
departure_date=data.get("departure_date"),
departure_time=data.get("departure_time"),
available_seats=_safe_int(data.get("available_seats")),
total_seats=_safe_int(data.get("total_seats")),
price=_safe_float(data.get("price")),
car_type=data.get("car_type"),
driver_name=data.get("driver_name"),
driver_phone=data.get("driver_phone"),
)
def _validate_extracted_trip(self, trip: ExtractedTrip) -> bool:
return all([
trip.departure,
trip.destination,
trip.departure_date,
trip.departure_time,
trip.driver_phone,
])
@staticmethod
def _normalize_phone(phone: str | None, country_code: str = "967") -> str | None:
if not phone:
return None
# Handle multiple phone numbers separated by /, or ,
phones = re.split(r'[/,]', phone)
normalized_phones: list[str] = []
for p in phones:
digits = "".join(c for c in p if c.isdigit())
if len(digits) >= 7:
digits = digits.lstrip("0")
if digits:
# Prepend country code if number doesn't already have it
if not digits.startswith(country_code):
digits = country_code + digits
normalized_phones.append(digits)
if not normalized_phones:
return None
# Return single phone or multiple phones separated by /
return "/".join(normalized_phones) if len(normalized_phones) > 1 else normalized_phones[0]
def _safe_int(value: Any) -> int | None:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _safe_float(value: Any) -> float | None:
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
|