codeBOKER commited on
Commit
bc2df26
·
1 Parent(s): 5a5bf7d

feat: support multiple phone numbers and prevent duplicate customers

Browse files

- Add support for multiple phone numbers separated by / in driver profiles
- Normalize phone numbers with country code prepending and leading zero stripping
- Prevent duplicate customers by checking phone_number before upserting
- Merge customer records when same phone appears in group vs DM context
- Make remote_jid nullable for unregistered drivers discovered via groups
- Add database migrations for nullable remote_jid and updated RPC

app/config.py CHANGED
@@ -47,6 +47,8 @@ class Settings(BaseSettings):
47
 
48
  admin_api_key: str = Field(min_length=1)
49
 
 
 
50
 
51
  @lru_cache
52
  def get_settings() -> Settings:
 
47
 
48
  admin_api_key: str = Field(min_length=1)
49
 
50
+ default_country_code: str = "967"
51
+
52
 
53
  @lru_cache
54
  def get_settings() -> Settings:
app/database/supabase.py CHANGED
@@ -35,7 +35,7 @@ class SupabaseRepository:
35
  async def upsert_customer(
36
  self,
37
  *,
38
- remote_jid: str,
39
  name: str | None = None,
40
  preferred_language: str | None = None,
41
  phone_number: str | None = None,
@@ -43,6 +43,33 @@ class SupabaseRepository:
43
  ) -> dict[str, Any]:
44
  if phone_number:
45
  phone_number = phone_number.split("@")[0]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
  payload = {
47
  "remoteJid": remote_jid,
48
  "name": name,
@@ -63,6 +90,7 @@ class SupabaseRepository:
63
  self,
64
  phone_number: str,
65
  ) -> dict[str, Any] | None:
 
66
  response = await (
67
  self.client.table("customers")
68
  .select("*")
@@ -70,7 +98,38 @@ class SupabaseRepository:
70
  .maybe_single()
71
  .execute()
72
  )
73
- return _response_data(response)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
 
75
  async def create_unregistered_driver_entities(
76
  self,
@@ -87,7 +146,7 @@ class SupabaseRepository:
87
  price: float,
88
  ) -> dict[str, Any]:
89
  customer = await self.upsert_customer(
90
- remote_jid=phone_number,
91
  name=driver_name,
92
  phone_number=phone_number,
93
  registered=False,
@@ -544,6 +603,36 @@ class SupabaseRepository:
544
  )
545
  return _response_data(response)
546
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
547
  async def create_driver(self, *, customer_id: str) -> dict[str, Any]:
548
  response = await (
549
  self.client.table("drivers")
 
35
  async def upsert_customer(
36
  self,
37
  *,
38
+ remote_jid: str | None = None,
39
  name: str | None = None,
40
  preferred_language: str | None = None,
41
  phone_number: str | None = None,
 
43
  ) -> dict[str, Any]:
44
  if phone_number:
45
  phone_number = phone_number.split("@")[0]
46
+
47
+ # Always check by phone_number first to avoid duplicates.
48
+ # A driver may have been saved from a group (remoteJid=None, phone_number=X),
49
+ # then later contacts us privately (remoteJid=Y, phone_number=X).
50
+ if phone_number:
51
+ existing = await self.get_customer_by_phone_number(phone_number)
52
+ if existing:
53
+ update_payload: dict[str, Any] = {}
54
+ if remote_jid and not existing.get("remoteJid"):
55
+ update_payload["remoteJid"] = remote_jid
56
+ if name is not None:
57
+ update_payload["name"] = name
58
+ if registered is not None:
59
+ update_payload["registered"] = registered
60
+ if phone_number is not None:
61
+ update_payload["phone_number"] = phone_number
62
+ if update_payload:
63
+ response = await (
64
+ self.client.table("customers")
65
+ .update(update_payload)
66
+ .eq("id", existing["id"])
67
+ .execute()
68
+ )
69
+ data = _response_data(response)
70
+ return data[0] if isinstance(data, list) else data
71
+ return existing
72
+
73
  payload = {
74
  "remoteJid": remote_jid,
75
  "name": name,
 
90
  self,
91
  phone_number: str,
92
  ) -> dict[str, Any] | None:
93
+ # Try exact match first
94
  response = await (
95
  self.client.table("customers")
96
  .select("*")
 
98
  .maybe_single()
99
  .execute()
100
  )
101
+ result = _response_data(response)
102
+ if result:
103
+ return result
104
+
105
+ # Try each individual phone from /-separated incoming number
106
+ phones = [p.strip() for p in phone_number.split("/")] if "/" in phone_number else [phone_number]
107
+ for phone in phones:
108
+ response = await (
109
+ self.client.table("customers")
110
+ .select("*")
111
+ .eq("phone_number", phone)
112
+ .maybe_single()
113
+ .execute()
114
+ )
115
+ result = _response_data(response)
116
+ if result:
117
+ return result
118
+
119
+ # Check if any stored customer has a /-separated phone_number containing our phone
120
+ for phone in phones:
121
+ response = await (
122
+ self.client.table("customers")
123
+ .select("*")
124
+ .like("phone_number", f"%{phone}%")
125
+ .maybe_single()
126
+ .execute()
127
+ )
128
+ result = _response_data(response)
129
+ if result:
130
+ return result
131
+
132
+ return None
133
 
134
  async def create_unregistered_driver_entities(
135
  self,
 
146
  price: float,
147
  ) -> dict[str, Any]:
148
  customer = await self.upsert_customer(
149
+ remote_jid=None,
150
  name=driver_name,
151
  phone_number=phone_number,
152
  registered=False,
 
603
  )
604
  return _response_data(response)
605
 
606
+ async def get_driver_by_phone_number(
607
+ self,
608
+ phone_number: str,
609
+ ) -> dict[str, Any] | None:
610
+ """Look up a driver by phone_number column.
611
+
612
+ Supports multiple phones separated by '/': tries each one until a match is found.
613
+ """
614
+ if "/" in phone_number:
615
+ phones = phone_number.split("/")
616
+ for phone in phones:
617
+ result = await self._get_driver_by_single_phone(phone.strip())
618
+ if result:
619
+ return result
620
+ return None
621
+ return await self._get_driver_by_single_phone(phone_number)
622
+
623
+ async def _get_driver_by_single_phone(
624
+ self,
625
+ phone: str,
626
+ ) -> dict[str, Any] | None:
627
+ response = await (
628
+ self.client.table("drivers")
629
+ .select("*, customers!inner(*)")
630
+ .eq("customers.phone_number", phone)
631
+ .maybe_single()
632
+ .execute()
633
+ )
634
+ return _response_data(response)
635
+
636
  async def create_driver(self, *, customer_id: str) -> dict[str, Any]:
637
  response = await (
638
  self.client.table("drivers")
app/services/group_message_service.py CHANGED
@@ -1,5 +1,6 @@
1
  import json
2
  import logging
 
3
  from datetime import date
4
  from pathlib import Path
5
  from typing import Any
@@ -71,7 +72,10 @@ class GroupMessageService:
71
  )
72
  return
73
 
74
- phone = self._normalize_phone(extracted.driver_phone)
 
 
 
75
  if not phone:
76
  logger.warning(
77
  "Group message %s: no valid phone number extracted",
@@ -90,7 +94,7 @@ class GroupMessageService:
90
  )
91
  return
92
 
93
- driver = await self.repository.get_driver_by_remoteJid(phone)
94
  if driver:
95
  existing_trip = await self.repository.get_driver_trip_by_datetime(
96
  driver_id=str(driver["id"]),
@@ -211,14 +215,29 @@ class GroupMessageService:
211
  ])
212
 
213
  @staticmethod
214
- def _normalize_phone(phone: str | None) -> str | None:
215
  if not phone:
216
  return None
217
- digits = "".join(c for c in phone if c.isdigit())
218
- if len(digits) < 7:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
219
  return None
220
- digits = digits.lstrip("0")
221
- return digits if digits else None
 
222
 
223
 
224
  def _safe_int(value: Any) -> int | None:
 
1
  import json
2
  import logging
3
+ import re
4
  from datetime import date
5
  from pathlib import Path
6
  from typing import Any
 
72
  )
73
  return
74
 
75
+ phone = self._normalize_phone(
76
+ extracted.driver_phone,
77
+ country_code=self.settings.default_country_code,
78
+ )
79
  if not phone:
80
  logger.warning(
81
  "Group message %s: no valid phone number extracted",
 
94
  )
95
  return
96
 
97
+ driver = await self.repository.get_driver_by_phone_number(phone)
98
  if driver:
99
  existing_trip = await self.repository.get_driver_trip_by_datetime(
100
  driver_id=str(driver["id"]),
 
215
  ])
216
 
217
  @staticmethod
218
+ def _normalize_phone(phone: str | None, country_code: str = "967") -> str | None:
219
  if not phone:
220
  return None
221
+
222
+ # Handle multiple phone numbers separated by /, or ,
223
+ phones = re.split(r'[/,]', phone)
224
+ normalized_phones: list[str] = []
225
+
226
+ for p in phones:
227
+ digits = "".join(c for c in p if c.isdigit())
228
+ if len(digits) >= 7:
229
+ digits = digits.lstrip("0")
230
+ if digits:
231
+ # Prepend country code if number doesn't already have it
232
+ if not digits.startswith(country_code):
233
+ digits = country_code + digits
234
+ normalized_phones.append(digits)
235
+
236
+ if not normalized_phones:
237
  return None
238
+
239
+ # Return single phone or multiple phones separated by /
240
+ return "/".join(normalized_phones) if len(normalized_phones) > 1 else normalized_phones[0]
241
 
242
 
243
  def _safe_int(value: Any) -> int | None:
app/tools/handlers.py CHANGED
@@ -264,8 +264,17 @@ class FalsaToolHandlers:
264
 
265
  driver_record = _first_or_dict(trip.get("drivers")) or {}
266
  driver_customer = driver_record.get("customers") or {}
267
- driver_recipient = driver_customer.get("phone_number") or driver_customer.get("remoteJid") or driver_record.get("remoteJid")
268
- driver_phone = driver_recipient.split("@")[0] if driver_recipient else None
 
 
 
 
 
 
 
 
 
269
 
270
  notification_status = "not_sent"
271
  notification_error = None
 
264
 
265
  driver_record = _first_or_dict(trip.get("drivers")) or {}
266
  driver_customer = driver_record.get("customers") or {}
267
+
268
+ # For unregistered drivers, phone_number may contain multiple numbers separated by /
269
+ driver_phone_raw = driver_customer.get("phone_number")
270
+ if driver_phone_raw and "/" in driver_phone_raw:
271
+ # Multiple phone numbers - return them separated by /
272
+ driver_phone = driver_phone_raw
273
+ driver_recipient = driver_phone_raw.split("/")[0]
274
+ else:
275
+ # Single phone number - use existing logic
276
+ driver_recipient = driver_phone_raw or driver_customer.get("remoteJid") or driver_record.get("remoteJid")
277
+ driver_phone = driver_recipient.split("@")[0] if driver_recipient else None
278
 
279
  notification_status = "not_sent"
280
  notification_error = None
prompts/group_trip_extraction.md CHANGED
@@ -26,7 +26,7 @@ If this IS a trip advertisement:
26
  "price": "number - price per seat",
27
  "car_type": "string or null - vehicle type if mentioned",
28
  "driver_name": "string or null - driver name if mentioned",
29
- "driver_phone": "string - phone number extracted from message text"
30
  }}
31
  ```
32
 
 
26
  "price": "number - price per seat",
27
  "car_type": "string or null - vehicle type if mentioned",
28
  "driver_name": "string or null - driver name if mentioned",
29
+ "driver_phone": "string - ALL phone numbers found in the message, separated by / if multiple. Extract from message text only."
30
  }}
31
  ```
32
 
scripts/run_migration.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Run the driver_message_card migration against Supabase.
3
+
4
+ Usage:
5
+ source .venv/bin/activate
6
+ python scripts/run_migration.py
7
+
8
+ Requires DATABASE_URL in .env or exports it.
9
+ The DATABASE_URL should be the direct connection string from:
10
+ Supabase Dashboard → Settings → Database → Connection string → URI (Session mode)
11
+ """
12
+
13
+ import os
14
+ import sys
15
+ from pathlib import Path
16
+
17
+ from dotenv import load_dotenv
18
+
19
+ load_dotenv()
20
+
21
+ try:
22
+ import psycopg
23
+ except ImportError:
24
+ sys.exit("psycopg not installed. Run: pip install 'psycopg[binary]'")
25
+
26
+ MIGRATION_SQL = (Path("supabase/migrations/202607140001_driver_message_card.sql")).read_text()
27
+
28
+
29
+ def main():
30
+ db_url = os.environ.get("DATABASE_URL")
31
+ if not db_url:
32
+ sys.exit(
33
+ "DATABASE_URL not set. Add it to .env or export it.\n"
34
+ "Get it from: Supabase Dashboard → Settings → Database → Connection string → URI"
35
+ )
36
+
37
+ print(f"Connecting to database...")
38
+ conn = psycopg.connect(db_url, connect_timeout=10)
39
+ try:
40
+ with conn.cursor() as cur:
41
+ cur.execute(MIGRATION_SQL)
42
+ conn.commit()
43
+ print("Migration applied successfully!")
44
+ except Exception as e:
45
+ conn.rollback()
46
+ sys.exit(f"Migration failed: {e}")
47
+ finally:
48
+ conn.close()
49
+
50
+
51
+ if __name__ == "__main__":
52
+ main()
supabase/migrations/202607160001_make_remote_jid_nullable.sql ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ -- Make remoteJid nullable so unregistered drivers from groups
2
+ -- can be created without a WhatsApp JID (they only have phone_number)
3
+ ALTER TABLE public.customers ALTER COLUMN "remoteJid" DROP NOT NULL;
4
+
5
+ -- Update unique constraint to allow NULL values (PostgreSQL allows multiple NULLs by default)
6
+ -- No additional action needed since UNIQUE already permits multiple NULLs
supabase/migrations/202607160002_update_rpc_phone_number.sql ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ -- Update match_active_trips RPC to use phone_number instead of remoteJid
2
+ -- This ensures unregistered drivers from groups (who have phone_number but no remoteJid)
3
+ -- are properly exposed in trip search results
4
+
5
+ DROP FUNCTION IF EXISTS public.match_active_trips(
6
+ extensions.vector(1024), float, int, text, text, date, text, time, int, text
7
+ );
8
+
9
+ CREATE OR REPLACE FUNCTION public.match_active_trips(
10
+ query_embedding extensions.vector(1024),
11
+ match_threshold float DEFAULT 0.0,
12
+ match_count int DEFAULT 10,
13
+ filter_departure text DEFAULT NULL,
14
+ filter_destination text DEFAULT NULL,
15
+ filter_departure_date date DEFAULT NULL,
16
+ filter_departure_time text DEFAULT NULL,
17
+ filter_requested_time time DEFAULT NULL,
18
+ filter_seats int DEFAULT 1,
19
+ filter_vehicle_type text DEFAULT NULL
20
+ )
21
+ RETURNS TABLE (
22
+ trip_id uuid,
23
+ departure text,
24
+ destination text,
25
+ departure_date date,
26
+ departure_time text,
27
+ available_seats integer,
28
+ total_seats integer,
29
+ price numeric,
30
+ status text,
31
+ driver_name text,
32
+ driver_phone_number text,
33
+ car_type text,
34
+ chunk_text text,
35
+ similarity float,
36
+ time_difference_minutes integer,
37
+ registered boolean
38
+ )
39
+ LANGUAGE sql STABLE
40
+ AS $$
41
+ WITH ranked AS (
42
+ SELECT
43
+ driver_trips.id AS trip_id,
44
+ driver_trips.departure,
45
+ driver_trips.destination,
46
+ driver_trips.departure_date,
47
+ driver_trips.departure_time,
48
+ driver_trips.available_seats,
49
+ driver_trips.total_seats,
50
+ driver_trips.price,
51
+ driver_trips.status,
52
+ customers.name AS driver_name,
53
+ COALESCE(customers.phone_number, customers."remoteJid") AS driver_phone_number,
54
+ driver_cars.car_type,
55
+ driver_trip_embeddings.chunk_text,
56
+ COALESCE(customers.registered, false) AS registered,
57
+ 1 - (driver_trip_embeddings.embedding <=> query_embedding) AS similarity,
58
+ driver_trip_embeddings.embedding <=> query_embedding AS vector_distance,
59
+ CASE
60
+ WHEN filter_requested_time IS NULL THEN NULL
61
+ ELSE abs(
62
+ extract(epoch FROM (
63
+ public.departure_bucket_clock_time(driver_trips.departure_time)
64
+ - filter_requested_time
65
+ )) / 60
66
+ )::integer
67
+ END AS time_difference_minutes
68
+ FROM public.driver_trip_embeddings
69
+ JOIN public.driver_trips ON driver_trips.id = driver_trip_embeddings.trip_id
70
+ LEFT JOIN public.drivers ON drivers.id = driver_trips.driver_id
71
+ LEFT JOIN public.customers ON customers.id = drivers.customer_id
72
+ LEFT JOIN public.driver_cars ON driver_cars.id = driver_trips.car_id
73
+ WHERE driver_trips.status = 'active'
74
+ AND driver_trips.available_seats >= COALESCE(filter_seats, 1)
75
+ AND (filter_departure IS NULL OR driver_trips.departure ILIKE '%' || filter_departure || '%')
76
+ AND (filter_destination IS NULL OR driver_trips.destination ILIKE '%' || filter_destination || '%')
77
+ AND (filter_departure_date IS NULL OR driver_trips.departure_date = filter_departure_date)
78
+ AND (filter_departure_time IS NULL OR driver_trips.departure_time = filter_departure_time)
79
+ AND (filter_vehicle_type IS NULL OR driver_cars.car_type ILIKE '%' || filter_vehicle_type || '%')
80
+ AND (
81
+ driver_trips.departure_date > (NOW() AT TIME ZONE 'Asia/Aden')::date
82
+ OR (
83
+ driver_trips.departure_date = (NOW() AT TIME ZONE 'Asia/Aden')::date
84
+ AND (
85
+ (NOW() AT TIME ZONE 'Asia/Aden')::time < TIME '12:00'
86
+ OR (
87
+ (NOW() AT TIME ZONE 'Asia/Aden')::time < TIME '18:00'
88
+ AND driver_trips.departure_time IN ('noon', 'night')
89
+ )
90
+ OR (
91
+ (NOW() AT TIME ZONE 'Asia/Aden')::time >= TIME '18:00'
92
+ AND driver_trips.departure_time = 'night'
93
+ )
94
+ )
95
+ )
96
+ )
97
+ )
98
+ SELECT
99
+ ranked.trip_id,
100
+ ranked.departure,
101
+ ranked.destination,
102
+ ranked.departure_date,
103
+ ranked.departure_time,
104
+ ranked.available_seats,
105
+ ranked.total_seats,
106
+ ranked.price,
107
+ ranked.status,
108
+ ranked.driver_name,
109
+ ranked.driver_phone_number,
110
+ ranked.car_type,
111
+ ranked.chunk_text,
112
+ ranked.similarity,
113
+ ranked.time_difference_minutes,
114
+ ranked.registered
115
+ FROM ranked
116
+ WHERE ranked.similarity >= match_threshold
117
+ ORDER BY
118
+ ranked.registered DESC,
119
+ ranked.time_difference_minutes NULLS LAST,
120
+ ranked.departure_date,
121
+ ranked.vector_distance
122
+ LIMIT match_count;
123
+ $$;
tests/conftest.py CHANGED
@@ -124,6 +124,20 @@ class FakeRepository:
124
  phone_number: str | None = None,
125
  registered: bool = True,
126
  ) -> dict[str, Any]:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
127
  jid = remote_jid or phone_number
128
  if not jid:
129
  raise ValueError("remote_jid or phone_number is required")
@@ -131,7 +145,7 @@ class FakeRepository:
131
  if customer is None:
132
  customer = {
133
  "id": f"cust-{len(self.customers_by_remote_jid) + 1}",
134
- "remoteJid": jid,
135
  "name": name,
136
  "preferred_language": preferred_language,
137
  "phone_number": phone_number,
@@ -178,7 +192,26 @@ class FakeRepository:
178
  self,
179
  phone_number: str,
180
  ) -> dict[str, Any] | None:
181
- return self.customers_by_phone.get(phone_number)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
182
 
183
  async def create_unregistered_driver_entities(
184
  self,
@@ -195,7 +228,7 @@ class FakeRepository:
195
  price: float,
196
  ) -> dict[str, Any]:
197
  customer = await self.upsert_customer(
198
- remote_jid=phone_number,
199
  name=driver_name,
200
  phone_number=phone_number,
201
  registered=False,
@@ -424,6 +457,36 @@ class FakeRepository:
424
  async def get_driver_by_remoteJid(self, remote_jid: str) -> dict[str, Any] | None:
425
  return self.drivers_by_remote_jid.get(remote_jid)
426
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
427
  async def create_driver(self, *, customer_id: str) -> dict[str, Any]:
428
  customer = next(
429
  (customer for customer in self.customers_by_remote_jid.values() if customer["id"] == customer_id),
 
124
  phone_number: str | None = None,
125
  registered: bool = True,
126
  ) -> dict[str, Any]:
127
+ # Always check by phone_number first to avoid duplicates
128
+ if phone_number:
129
+ existing = await self.get_customer_by_phone_number(phone_number)
130
+ if existing:
131
+ if remote_jid and not existing.get("remoteJid"):
132
+ existing["remoteJid"] = remote_jid
133
+ if name is not None:
134
+ existing["name"] = name
135
+ if registered is not None:
136
+ existing["registered"] = registered
137
+ if phone_number is not None:
138
+ existing["phone_number"] = phone_number
139
+ return existing
140
+
141
  jid = remote_jid or phone_number
142
  if not jid:
143
  raise ValueError("remote_jid or phone_number is required")
 
145
  if customer is None:
146
  customer = {
147
  "id": f"cust-{len(self.customers_by_remote_jid) + 1}",
148
+ "remoteJid": remote_jid,
149
  "name": name,
150
  "preferred_language": preferred_language,
151
  "phone_number": phone_number,
 
192
  self,
193
  phone_number: str,
194
  ) -> dict[str, Any] | None:
195
+ # Try exact match first
196
+ result = self.customers_by_phone.get(phone_number)
197
+ if result:
198
+ return result
199
+
200
+ # Try each individual phone from /-separated incoming number
201
+ phones = [p.strip() for p in phone_number.split("/")] if "/" in phone_number else [phone_number]
202
+ for phone in phones:
203
+ result = self.customers_by_phone.get(phone)
204
+ if result:
205
+ return result
206
+
207
+ # Check if any stored customer has a /-separated phone_number containing our phone
208
+ for customer in self.customers_by_remote_jid.values():
209
+ stored = customer.get("phone_number") or ""
210
+ for phone in phones:
211
+ if phone in stored.split("/"):
212
+ return customer
213
+
214
+ return None
215
 
216
  async def create_unregistered_driver_entities(
217
  self,
 
228
  price: float,
229
  ) -> dict[str, Any]:
230
  customer = await self.upsert_customer(
231
+ remote_jid=None,
232
  name=driver_name,
233
  phone_number=phone_number,
234
  registered=False,
 
457
  async def get_driver_by_remoteJid(self, remote_jid: str) -> dict[str, Any] | None:
458
  return self.drivers_by_remote_jid.get(remote_jid)
459
 
460
+ async def get_driver_by_phone_number(
461
+ self,
462
+ phone_number: str,
463
+ ) -> dict[str, Any] | None:
464
+ """Look up a driver by phone_number column.
465
+
466
+ Supports multiple phones separated by '/': tries each one until a match is found.
467
+ """
468
+ if "/" in phone_number:
469
+ phones = phone_number.split("/")
470
+ for phone in phones:
471
+ result = await self._get_driver_by_single_phone(phone.strip())
472
+ if result:
473
+ return result
474
+ return None
475
+ return await self._get_driver_by_single_phone(phone_number)
476
+
477
+ async def _get_driver_by_single_phone(
478
+ self,
479
+ phone: str,
480
+ ) -> dict[str, Any] | None:
481
+ for customer in self.customers_by_remote_jid.values():
482
+ stored_phone = customer.get("phone_number") or ""
483
+ # Match exact or within /-separated list
484
+ if phone == stored_phone or phone in stored_phone.split("/"):
485
+ for driver in self.drivers_by_remote_jid.values():
486
+ if driver.get("customer_id") == customer["id"]:
487
+ return driver
488
+ return None
489
+
490
  async def create_driver(self, *, customer_id: str) -> dict[str, Any]:
491
  customer = next(
492
  (customer for customer in self.customers_by_remote_jid.values() if customer["id"] == customer_id),
tests/test_group_message_service.py CHANGED
@@ -415,3 +415,257 @@ async def test_trip_ad_with_missing_name_uses_none(settings: Settings) -> None:
415
  assert len(repo.created_trips) == 1
416
  customer = list(repo.customers_by_remote_jid.values())[0]
417
  assert customer["name"] is None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
415
  assert len(repo.created_trips) == 1
416
  customer = list(repo.customers_by_remote_jid.values())[0]
417
  assert customer["name"] is None
418
+
419
+
420
+ @pytest.mark.asyncio
421
+ async def test_multiple_phone_numbers_are_separated_by_slash(settings: Settings) -> None:
422
+ repo = FakeRepository()
423
+ embeddings = FakeEmbeddings()
424
+ provider = FakeProvider(
425
+ response_content=_trip_ad_response(driver_phone="967712345678/967876543210")
426
+ )
427
+
428
+ service = GroupMessageService(
429
+ repository=repo,
430
+ embeddings=embeddings,
431
+ ai=provider,
432
+ settings=settings,
433
+ )
434
+
435
+ inbound = _make_inbound(text="رحلة")
436
+ await service.handle_group_message(inbound)
437
+
438
+ assert len(repo.customers_by_remote_jid) == 1
439
+ customer = list(repo.customers_by_remote_jid.values())[0]
440
+ assert customer["phone_number"] == "967712345678/967876543210"
441
+
442
+
443
+ @pytest.mark.asyncio
444
+ async def test_unregistered_driver_has_null_remote_jid(settings: Settings) -> None:
445
+ repo = FakeRepository()
446
+ embeddings = FakeEmbeddings()
447
+ provider = FakeProvider(response_content=_trip_ad_response())
448
+
449
+ service = GroupMessageService(
450
+ repository=repo,
451
+ embeddings=embeddings,
452
+ ai=provider,
453
+ settings=settings,
454
+ )
455
+
456
+ inbound = _make_inbound(text="رحلة")
457
+ await service.handle_group_message(inbound)
458
+
459
+ assert len(repo.customers_by_remote_jid) == 1
460
+ customer = list(repo.customers_by_remote_jid.values())[0]
461
+ assert customer["remoteJid"] is None
462
+ assert customer["phone_number"] == "967712345678"
463
+
464
+
465
+ @pytest.mark.asyncio
466
+ async def test_phone_normalization_with_multiple_numbers(settings: Settings) -> None:
467
+ repo = FakeRepository()
468
+ embeddings = FakeEmbeddings()
469
+ provider = FakeProvider(
470
+ response_content=_trip_ad_response(driver_phone="+967-71-234-5678 / +967-78-765-4321")
471
+ )
472
+
473
+ service = GroupMessageService(
474
+ repository=repo,
475
+ embeddings=embeddings,
476
+ ai=provider,
477
+ settings=settings,
478
+ )
479
+
480
+ inbound = _make_inbound(text="رحلة")
481
+ await service.handle_group_message(inbound)
482
+
483
+ assert len(repo.customers_by_remote_jid) == 1
484
+ customer = list(repo.customers_by_remote_jid.values())[0]
485
+ assert customer["phone_number"] == "967712345678/967787654321"
486
+
487
+
488
+ @pytest.mark.asyncio
489
+ async def test_local_phone_number_gets_country_code_prepended(settings: Settings) -> None:
490
+ repo = FakeRepository()
491
+ embeddings = FakeEmbeddings()
492
+ provider = FakeProvider(
493
+ response_content=_trip_ad_response(driver_phone="712345678")
494
+ )
495
+
496
+ service = GroupMessageService(
497
+ repository=repo,
498
+ embeddings=embeddings,
499
+ ai=provider,
500
+ settings=settings,
501
+ )
502
+
503
+ inbound = _make_inbound(text="رحلة")
504
+ await service.handle_group_message(inbound)
505
+
506
+ assert len(repo.customers_by_remote_jid) == 1
507
+ customer = list(repo.customers_by_remote_jid.values())[0]
508
+ assert customer["phone_number"] == "967712345678"
509
+
510
+
511
+ @pytest.mark.asyncio
512
+ async def test_leading_zero_stripped_then_country_code_prepended(settings: Settings) -> None:
513
+ repo = FakeRepository()
514
+ embeddings = FakeEmbeddings()
515
+ provider = FakeProvider(
516
+ response_content=_trip_ad_response(driver_phone="0712345678")
517
+ )
518
+
519
+ service = GroupMessageService(
520
+ repository=repo,
521
+ embeddings=embeddings,
522
+ ai=provider,
523
+ settings=settings,
524
+ )
525
+
526
+ inbound = _make_inbound(text="رحلة")
527
+ await service.handle_group_message(inbound)
528
+
529
+ assert len(repo.customers_by_remote_jid) == 1
530
+ customer = list(repo.customers_by_remote_jid.values())[0]
531
+ assert customer["phone_number"] == "967712345678"
532
+
533
+
534
+ @pytest.mark.asyncio
535
+ async def test_existing_unregistered_driver_with_multiple_phones_adds_new_trip(
536
+ settings: Settings,
537
+ ) -> None:
538
+ repo = FakeRepository()
539
+ embeddings = FakeEmbeddings()
540
+ provider = FakeProvider(
541
+ response_content=_trip_ad_response(driver_phone="967712345678/967876543210")
542
+ )
543
+
544
+ await repo.upsert_customer(
545
+ remote_jid=None,
546
+ name="أحمد",
547
+ phone_number="967712345678/967876543210",
548
+ registered=False,
549
+ )
550
+
551
+ service = GroupMessageService(
552
+ repository=repo,
553
+ embeddings=embeddings,
554
+ ai=provider,
555
+ settings=settings,
556
+ )
557
+
558
+ inbound = _make_inbound(text="رحلة من صنعاء إلى عدن")
559
+ await service.handle_group_message(inbound)
560
+
561
+ assert len(repo.created_trips) == 1
562
+ assert repo.created_trips[0]["departure"] == "صنعاء"
563
+
564
+
565
+ @pytest.mark.asyncio
566
+ async def test_get_driver_by_phone_number_with_multiple_phones(settings: Settings) -> None:
567
+ repo = FakeRepository()
568
+
569
+ customer = await repo.upsert_customer(
570
+ remote_jid=None,
571
+ name="أحمد",
572
+ phone_number="967712345678/967876543210",
573
+ registered=False,
574
+ )
575
+ await repo.create_driver(customer_id=str(customer["id"]))
576
+
577
+ driver = await repo.get_driver_by_phone_number("967712345678/967876543210")
578
+ assert driver is not None
579
+ assert driver["customers"]["phone_number"] == "967712345678/967876543210"
580
+
581
+
582
+ @pytest.mark.asyncio
583
+ async def test_unregistered_driver_later_dm_merges_into_same_row(settings: Settings) -> None:
584
+ """When a driver saved from a group (remoteJid=None) later DMs us,
585
+ the DM should update the existing row instead of creating a duplicate."""
586
+ repo = FakeRepository()
587
+ embeddings = FakeEmbeddings()
588
+ provider = FakeProvider(response_content=_trip_ad_response())
589
+
590
+ # Step 1: Driver posts trip ad in group (creates unregistered customer)
591
+ service = GroupMessageService(
592
+ repository=repo,
593
+ embeddings=embeddings,
594
+ ai=provider,
595
+ settings=settings,
596
+ )
597
+ inbound = _make_inbound(text="رحلة من صنعاء إلى عدن")
598
+ await service.handle_group_message(inbound)
599
+
600
+ assert len(repo.customers_by_remote_jid) == 1
601
+ existing = list(repo.customers_by_remote_jid.values())[0]
602
+ assert existing["remoteJid"] is None
603
+ assert existing["phone_number"] == "967712345678"
604
+ existing_id = existing["id"]
605
+
606
+ # Step 2: Same driver sends a DM (should update existing row, not create new one)
607
+ customer = await repo.upsert_customer(
608
+ remote_jid="967712345678",
609
+ name="أحمد",
610
+ phone_number="967712345678",
611
+ registered=True,
612
+ )
613
+
614
+ assert customer["id"] == existing_id
615
+ assert customer["remoteJid"] == "967712345678"
616
+ assert customer["registered"] is True
617
+ # Should still be only one customer row
618
+ assert len(repo.customers_by_remote_jid) == 1
619
+
620
+
621
+ @pytest.mark.asyncio
622
+ async def test_dm_with_shared_phone_merges_into_existing(settings: Settings) -> None:
623
+ """When a customer exists with phone_number matching a DM's phone,
624
+ the DM should update the existing row."""
625
+ repo = FakeRepository()
626
+
627
+ # Create a customer via DM first
628
+ customer1 = await repo.upsert_customer(
629
+ remote_jid="967712345678",
630
+ name="أحمد",
631
+ phone_number="967712345678",
632
+ registered=True,
633
+ )
634
+
635
+ # Another DM with same phone_number but different remoteJid (e.g. new device)
636
+ customer2 = await repo.upsert_customer(
637
+ remote_jid="967999999999",
638
+ name="أحمد",
639
+ phone_number="967712345678",
640
+ registered=True,
641
+ )
642
+
643
+ assert customer1["id"] == customer2["id"]
644
+ assert len(repo.customers_by_remote_jid) == 1
645
+
646
+
647
+ @pytest.mark.asyncio
648
+ async def test_dm_with_multi_phone_merges_into_existing(settings: Settings) -> None:
649
+ """When a DM's phone_number matches one of the phones in an existing
650
+ /-separated phone_number, it should update the existing row."""
651
+ repo = FakeRepository()
652
+
653
+ # Create unregistered driver with two phones
654
+ customer1 = await repo.upsert_customer(
655
+ remote_jid=None,
656
+ name="أحمد",
657
+ phone_number="967712345678/967876543210",
658
+ registered=False,
659
+ )
660
+
661
+ # DM comes in with just the first phone number
662
+ customer2 = await repo.upsert_customer(
663
+ remote_jid="967712345678",
664
+ name="أحمد",
665
+ phone_number="967712345678",
666
+ registered=True,
667
+ )
668
+
669
+ assert customer1["id"] == customer2["id"]
670
+ assert customer2["remoteJid"] == "967712345678"
671
+ assert customer2["registered"] is True