296 lines
12 KiB
Python
Executable File
296 lines
12 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
"""Mock AI Agent — Asterisk/ARI 版.
|
|
外呼流程 (控制面纯 REST + RTP; 另加一个原生 socket websocket 仅用于保持 Stasis app 运行,
|
|
否则 ARI 会把进入 app 的通道直接挂断):
|
|
0. 连接 /ari/events?app=outbound —— 让 Stasis app 处于运行态 (只丢事件, 不解析业务)
|
|
1. POST /bridges 建立 mixing 桥
|
|
2. POST /channels/externalMedia 外部媒体通道 (ulaw), Asterisk 把桥内混音发到 external_host,
|
|
本侧回发地址 = 通道变量 UNICASTRTP_LOCAL_ADDRESS/PORT (兜底: 首包源地址)
|
|
3. POST /channels originate 经 PJSIP trunk 外呼被叫
|
|
4. 轮询被叫通道至 Up (应答), 双通道入桥
|
|
5. RTP 双向: RX=被叫音频(模拟 ASR 输入, RMS 统计); TX=440Hz 间歇音(模拟 LLM/TTS 输出)
|
|
6. DELETE 通道+桥 → BYE
|
|
纯 stdlib (urllib + asyncio + 原生 socket)。换真实 LLM/ASR: 替换 TX 音源与 RX 消费即可。
|
|
用法: mock-agent.py --ari-base http://asterisk1:8088 --external-host <本容器名>:40001 \
|
|
--number +15105550123 [--duration 15]
|
|
"""
|
|
import argparse, asyncio, base64, json, math, os, random, socket, struct, sys, threading, time
|
|
import urllib.error, urllib.parse, urllib.request
|
|
|
|
def log(*a): print(f"[agent {time.strftime('%H:%M:%S')}]", *a, flush=True)
|
|
|
|
# ---------- G.711 u-law (与 mock-provider 同源) ----------
|
|
def ulaw_encode(s: int) -> int:
|
|
s = max(-32635, min(32635, s)); sign = 0x80 if s < 0 else 0
|
|
if sign: s = -s
|
|
s += 0x84
|
|
e = 7
|
|
for i in range(7, -1, -1):
|
|
if s & (0x40 << i): e = i; break
|
|
return (~(sign | (e << 4) | ((s >> (e + 3)) & 0x0F))) & 0xFF
|
|
|
|
def ulaw_decode(u: int) -> int:
|
|
u = ~u & 0xFF
|
|
sign = u & 0x80; e = (u >> 4) & 0x07; m = u & 0x0F
|
|
s = ((m << (e + 3)) + 0x84) << 2
|
|
return -s if sign else s
|
|
|
|
def tone_payload(ts: float) -> bytes:
|
|
"""20ms / 160 samples @8kHz, 440Hz 500ms-on 500ms-off (模拟 TTS)"""
|
|
out = bytearray()
|
|
on = (ts % 1.0) < 0.5
|
|
for i in range(160):
|
|
t = ts + i / 8000.0
|
|
v = int(6000 * math.sin(2 * math.pi * 440 * t)) if on else 0
|
|
out.append(ulaw_encode(v))
|
|
return bytes(out)
|
|
|
|
# ---------- ARI REST ----------
|
|
class Ari:
|
|
def __init__(self, base, user, password):
|
|
self.base = base.rstrip("/")
|
|
self.auth = "Basic " + base64.b64encode(f"{user}:{password}".encode()).decode()
|
|
|
|
def req(self, method, path, params=None):
|
|
url = f"{self.base}/ari{path}"
|
|
if params:
|
|
url += "?" + urllib.parse.urlencode(params)
|
|
r = urllib.request.Request(url, method=method)
|
|
r.add_header("Authorization", self.auth)
|
|
try:
|
|
with urllib.request.urlopen(r, timeout=15) as resp:
|
|
raw = resp.read()
|
|
return resp.status, (json.loads(raw) if raw else None)
|
|
except urllib.error.HTTPError as e:
|
|
raw = e.read()
|
|
try:
|
|
return e.code, (json.loads(raw) if raw else None)
|
|
except Exception:
|
|
return e.code, {"error": raw.decode("utf-8", "replace")}
|
|
|
|
# ---------- ARI Stasis app 保活 (极简事件 websocket, 纯 stdlib) ----------
|
|
class AppKeeper:
|
|
"""ARI 要求 Stasis app 有活跃 websocket 订阅, 否则进入 app 的通道会被立即挂断.
|
|
本类只保持连接并丢弃事件, 控制面仍走 REST."""
|
|
|
|
def __init__(self, base, app, user, password):
|
|
u = urllib.parse.urlparse(base)
|
|
self.host, self.port = u.hostname, u.port or 80
|
|
self.app, self.user, self.password = app, user, password
|
|
self.sock = None
|
|
|
|
def start(self):
|
|
key = base64.b64encode(os.urandom(16)).decode()
|
|
req = ("GET /ari/events?app=" + urllib.parse.quote(self.app)
|
|
+ "&api_key=" + urllib.parse.quote(f"{self.user}:{self.password}")
|
|
+ f" HTTP/1.1\r\nHost: {self.host}:{self.port}\r\n"
|
|
+ "Upgrade: websocket\r\nConnection: Upgrade\r\n"
|
|
+ f"Sec-WebSocket-Key: {key}\r\nSec-WebSocket-Version: 13\r\n\r\n")
|
|
self.sock = socket.create_connection((self.host, self.port), timeout=10)
|
|
self.sock.sendall(req.encode())
|
|
resp = b""
|
|
while b"\r\n\r\n" not in resp:
|
|
chunk = self.sock.recv(1) # 逐字节读响应头, 不吞后续 ws 帧
|
|
if not chunk:
|
|
raise ConnectionError("ws closed during handshake")
|
|
resp += chunk
|
|
if b"101" not in resp.split(b"\r\n", 1)[0]:
|
|
raise ConnectionError(f"ws handshake failed: {resp[:120]!r}")
|
|
self.sock.settimeout(None)
|
|
threading.Thread(target=self._drain, daemon=True).start()
|
|
|
|
def _recv_exact(self, n):
|
|
buf = b""
|
|
while len(buf) < n:
|
|
c = self.sock.recv(n - len(buf))
|
|
if not c:
|
|
raise ConnectionError("ws closed")
|
|
buf += c
|
|
return buf
|
|
|
|
def _drain(self):
|
|
try:
|
|
while True:
|
|
b0 = self._recv_exact(1)[0]
|
|
b1 = self._recv_exact(1)[0]
|
|
ln = b1 & 0x7F
|
|
if ln == 126:
|
|
ln = struct.unpack("!H", self._recv_exact(2))[0]
|
|
elif ln == 127:
|
|
ln = struct.unpack("!Q", self._recv_exact(8))[0]
|
|
payload = self._recv_exact(ln) if ln else b""
|
|
op = b0 & 0x0F
|
|
if op == 0x8: # close
|
|
break
|
|
if op == 0x9: # ping → pong
|
|
mask = os.urandom(4) # client 帧必须带 mask
|
|
masked = bytes(b ^ mask[i % 4] for i, b in enumerate(payload))
|
|
self.sock.sendall(bytes((0x8A, 0x80 | (len(payload) & 0x7F))) + mask + masked)
|
|
continue
|
|
if op == 0x1 and payload:
|
|
try:
|
|
ev = json.loads(payload)
|
|
except Exception:
|
|
continue
|
|
if ev.get("type") in ("StasisStart", "StasisEnd", "ChannelDestroyed"):
|
|
log(f"event {ev['type']} {(ev.get('channel') or {}).get('name', '')}")
|
|
except Exception:
|
|
pass
|
|
|
|
def stop(self):
|
|
try:
|
|
if self.sock:
|
|
self.sock.close()
|
|
except Exception:
|
|
pass
|
|
|
|
# ---------- RTP 媒体 ----------
|
|
class Rx(asyncio.DatagramProtocol):
|
|
def __init__(self, stats):
|
|
self.stats, self.transport = stats, None
|
|
|
|
def connection_made(self, transport):
|
|
self.transport = transport
|
|
|
|
def datagram_received(self, data, addr):
|
|
if len(data) < 12: return
|
|
self.stats["rx"] += 1
|
|
payload = data[12:]
|
|
dec = [ulaw_decode(b) for b in payload[:40]]
|
|
self.stats["amp_sum"] += sum(abs(x) for x in dec) / max(1, len(dec))
|
|
if self.stats["tx_to"] is None: # 兜底: 用首包源地址作回程
|
|
self.stats["tx_to"] = addr
|
|
if self.stats["rx"] % 250 == 0:
|
|
log(f"RX frames={self.stats['rx']} avg_rms={self.stats['amp_sum']/self.stats['rx']:.0f}")
|
|
|
|
async def media_loop(listen_port, tx_to, duration):
|
|
stats = {"rx": 0, "amp_sum": 0.0, "tx": 0, "tx_to": tx_to}
|
|
loop = asyncio.get_running_loop()
|
|
rx_t, _ = await loop.create_datagram_endpoint(
|
|
lambda: Rx(stats), local_addr=("0.0.0.0", listen_port))
|
|
|
|
async def tx():
|
|
seq = random.randint(0, 65535); ts = random.randint(0, 10**6)
|
|
ssrc = random.randint(1, 2**31)
|
|
t0 = time.time()
|
|
while time.time() - t0 < duration:
|
|
if stats["tx_to"]:
|
|
payload = tone_payload(time.time() - t0)
|
|
hdr = struct.pack("!BBHII", 0x80, 0, seq & 0xFFFF, ts & 0xFFFFFFFF, ssrc)
|
|
rx_t.sendto(hdr + payload, stats["tx_to"])
|
|
stats["tx"] += 1
|
|
seq += 1; ts += 160
|
|
await asyncio.sleep(0.02)
|
|
|
|
try:
|
|
await asyncio.wait_for(tx(), timeout=duration + 5)
|
|
finally:
|
|
rx_t.close()
|
|
return stats
|
|
|
|
# ---------- 主流程 ----------
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--ari-base", default="http://127.0.0.1:8088")
|
|
ap.add_argument("--ari-user", default="outbound")
|
|
ap.add_argument("--ari-pass", default="ari_mock_7c3f9e1b5a2d4806")
|
|
ap.add_argument("--app", default="outbound")
|
|
ap.add_argument("--external-host", required=True, help="本侧 RTP 地址 ip:port (供 Asterisk 回发混音)")
|
|
ap.add_argument("--number", default="+15105550123")
|
|
ap.add_argument("--caller-id", default="+15109990001")
|
|
ap.add_argument("--trunk", default="mock-trunk")
|
|
ap.add_argument("--duration", type=float, default=15)
|
|
ap.add_argument("--ring-timeout", type=float, default=30)
|
|
args = ap.parse_args()
|
|
|
|
listen_port = int(args.external_host.rsplit(":", 1)[1])
|
|
ari = Ari(args.ari_base, args.ari_user, args.ari_pass)
|
|
suffix = f"{int(time.time())}-{random.randint(100, 999)}"
|
|
bridge_id, em_id, callee_id = f"br-{suffix}", f"em-{suffix}", f"callee-{suffix}"
|
|
|
|
# 0. 保持 Stasis app 运行 (否则通道进 app 即被挂断)
|
|
keeper = AppKeeper(args.ari_base, args.app, args.ari_user, args.ari_pass)
|
|
keeper.start()
|
|
log(f"stasis app '{args.app}' running (event ws connected)")
|
|
|
|
def fail(msg, detail=None):
|
|
log(f"FAIL: {msg}", detail or "")
|
|
cleanup()
|
|
sys.exit(1)
|
|
|
|
def cleanup():
|
|
for m, p in (("DELETE", f"/channels/{callee_id}"), ("DELETE", f"/channels/{em_id}"),
|
|
("DELETE", f"/bridges/{bridge_id}")):
|
|
try: ari.req(m, p)
|
|
except Exception: pass
|
|
keeper.stop()
|
|
|
|
# 1. 混音桥
|
|
st, r = ari.req("POST", "/bridges", params={"type": "mixing", "bridgeId": bridge_id})
|
|
if st not in (200, 201): fail("create bridge", (st, r))
|
|
|
|
# 2. 外部媒体通道 (Asterisk → 本侧 RTP)
|
|
st, em = ari.req("POST", "/channels/externalMedia", params={
|
|
"app": args.app, "external_host": args.external_host,
|
|
"format": "ulaw", "channelId": em_id})
|
|
if st not in (200, 201): fail("create externalMedia", (st, r))
|
|
log(f"externalMedia channel up: {em.get('name')}")
|
|
|
|
# 3. Asterisk 侧 RTP 回送地址
|
|
tx_to = None
|
|
st, v = ari.req("GET", f"/channels/{em_id}/variable", {"variable": "UNICASTRTP_LOCAL_ADDRESS"})
|
|
ip = (v or {}).get("value") if st == 200 else None
|
|
st, v = ari.req("GET", f"/channels/{em_id}/variable", {"variable": "UNICASTRTP_LOCAL_PORT"})
|
|
port = (v or {}).get("value") if st == 200 else None
|
|
if ip and port:
|
|
tx_to = (ip, int(port))
|
|
log(f"asterisk RTP return addr: {ip}:{port}")
|
|
else:
|
|
log("UNICASTRTP vars not set; will reply-to-source on first RX")
|
|
|
|
# 4. 发起外呼
|
|
st, ch = ari.req("POST", "/channels", params={
|
|
"endpoint": f"PJSIP/{args.number}@{args.trunk}",
|
|
"app": args.app, "appArgs": "outbound",
|
|
"callerId": args.caller_id, "channelId": callee_id,
|
|
"timeout": int(args.ring_timeout)})
|
|
if st not in (200, 201): fail("originate", (st, ch))
|
|
log(f"INVITE sent via {args.trunk} -> {args.number} (channel {callee_id})")
|
|
|
|
# 5. 等待被叫应答
|
|
deadline = time.time() + args.ring_timeout
|
|
state = None
|
|
while time.time() < deadline:
|
|
st, ch = ari.req("GET", f"/channels/{callee_id}")
|
|
if st == 200:
|
|
state = ch.get("state")
|
|
elif st in (404, 410):
|
|
fail("callee channel gone before answer")
|
|
if state == "Up":
|
|
break
|
|
time.sleep(0.25)
|
|
if state != "Up":
|
|
fail(f"no answer (state={state})")
|
|
log(f"callee answered: state={state}")
|
|
|
|
# 6. 双通道入桥 → 双向 RTP
|
|
st, r = ari.req("POST", f"/bridges/{bridge_id}/addChannel",
|
|
params={"channel": f"{em_id},{callee_id}"})
|
|
if st not in (200, 201, 204): fail("addChannel", (st, r))
|
|
log(f"bridged [{em_id} + {callee_id}], streaming {args.duration}s "
|
|
f"(RX=ASR mock, TX=440Hz LLM/TTS mock)")
|
|
|
|
# 7. 媒体收发
|
|
stats = asyncio.run(media_loop(listen_port, tx_to, args.duration))
|
|
|
|
# 8. 挂断收尾
|
|
cleanup()
|
|
n = stats["rx"]
|
|
amp = stats["amp_sum"] / n if n else 0.0
|
|
print(f"RESULT number={args.number} rx_frames={n} rx_avg_rms={amp:.0f} "
|
|
f"tx_frames={stats['tx']} bidirectional={'YES' if n > 50 and amp > 100 else 'NO'}",
|
|
flush=True)
|
|
|
|
if __name__ == "__main__":
|
|
main()
|