openspline 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- openspline-0.1.0/.gitignore +21 -0
- openspline-0.1.0/LICENSE +21 -0
- openspline-0.1.0/PKG-INFO +12 -0
- openspline-0.1.0/pyproject.toml +14 -0
- openspline-0.1.0/src/openspline/__init__.py +20 -0
- openspline-0.1.0/src/openspline/client.py +301 -0
- openspline-0.1.0/src/openspline/config.py +49 -0
- openspline-0.1.0/src/openspline/errors.py +30 -0
- openspline-0.1.0/src/openspline/integrations/__init__.py +6 -0
- openspline-0.1.0/src/openspline/integrations/elevenlabs.py +72 -0
- openspline-0.1.0/src/openspline/integrations/events.py +153 -0
- openspline-0.1.0/src/openspline/integrations/livekit.py +68 -0
- openspline-0.1.0/src/openspline/integrations/pipecat.py +87 -0
- openspline-0.1.0/src/openspline/integrations/ten.py +45 -0
- openspline-0.1.0/src/openspline/integrations/vocode.py +83 -0
- openspline-0.1.0/src/openspline/media.py +88 -0
- openspline-0.1.0/src/openspline/py.typed +0 -0
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
.venv/
|
|
2
|
+
__pycache__/
|
|
3
|
+
*.py[cod]
|
|
4
|
+
*.egg-info/
|
|
5
|
+
.env
|
|
6
|
+
.env.*
|
|
7
|
+
!.env.example
|
|
8
|
+
node_modules/
|
|
9
|
+
dist/
|
|
10
|
+
.pytest_cache/
|
|
11
|
+
.ruff_cache/
|
|
12
|
+
/models/
|
|
13
|
+
/runtime/
|
|
14
|
+
artifacts/private/
|
|
15
|
+
*.log
|
|
16
|
+
.gstack/
|
|
17
|
+
|
|
18
|
+
.venvs/
|
|
19
|
+
.cache/
|
|
20
|
+
.tools/
|
|
21
|
+
experiments/
|
openspline-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 openspline contributors
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: openspline
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Add a live face to any voice agent
|
|
5
|
+
License-Expression: MIT
|
|
6
|
+
License-File: LICENSE
|
|
7
|
+
Requires-Python: <3.13,>=3.10
|
|
8
|
+
Requires-Dist: httpx<1,>=0.27
|
|
9
|
+
Requires-Dist: websockets<16,>=12
|
|
10
|
+
Provides-Extra: media
|
|
11
|
+
Requires-Dist: aiortc<2,>=1.14; extra == 'media'
|
|
12
|
+
Requires-Dist: av<17,>=14; extra == 'media'
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["hatchling"]
|
|
3
|
+
build-backend = "hatchling.build"
|
|
4
|
+
[project]
|
|
5
|
+
name = "openspline"
|
|
6
|
+
version = "0.1.0"
|
|
7
|
+
description = "Add a live face to any voice agent"
|
|
8
|
+
requires-python = ">=3.10,<3.13"
|
|
9
|
+
license = "MIT"
|
|
10
|
+
dependencies = ["httpx>=0.27,<1", "websockets>=12,<16"]
|
|
11
|
+
[project.optional-dependencies]
|
|
12
|
+
media = ["aiortc>=1.14,<2", "av>=14,<17"]
|
|
13
|
+
[tool.hatch.build.targets.wheel]
|
|
14
|
+
packages = ["src/openspline"]
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
from .client import AvatarSession, Openspline
|
|
2
|
+
from .config import OpensplineConfig
|
|
3
|
+
from .errors import (
|
|
4
|
+
CapacityError,
|
|
5
|
+
ConfigurationError,
|
|
6
|
+
ConnectionError,
|
|
7
|
+
InferenceError,
|
|
8
|
+
OpensplineError,
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
__all__ = [
|
|
12
|
+
"AvatarSession",
|
|
13
|
+
"CapacityError",
|
|
14
|
+
"ConfigurationError",
|
|
15
|
+
"ConnectionError",
|
|
16
|
+
"InferenceError",
|
|
17
|
+
"Openspline",
|
|
18
|
+
"OpensplineConfig",
|
|
19
|
+
"OpensplineError",
|
|
20
|
+
]
|
|
@@ -0,0 +1,301 @@
|
|
|
1
|
+
"""Lightweight async client. No CUDA, NumPy, or framework imports."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import base64
|
|
7
|
+
import contextlib
|
|
8
|
+
import inspect
|
|
9
|
+
import json
|
|
10
|
+
import uuid
|
|
11
|
+
from collections.abc import Mapping
|
|
12
|
+
from dataclasses import fields
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
|
|
15
|
+
import httpx
|
|
16
|
+
from websockets.legacy.client import connect
|
|
17
|
+
|
|
18
|
+
from .config import OpensplineConfig, Quality
|
|
19
|
+
from .errors import ConnectionError, error_from
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class Openspline:
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
url: str | None = None,
|
|
26
|
+
timeout: float | None = None,
|
|
27
|
+
*,
|
|
28
|
+
quality: Quality | None = None,
|
|
29
|
+
viewer_timeout: float | None = None,
|
|
30
|
+
config: OpensplineConfig | Mapping | None = None,
|
|
31
|
+
):
|
|
32
|
+
if config is None:
|
|
33
|
+
values = {}
|
|
34
|
+
elif isinstance(config, OpensplineConfig):
|
|
35
|
+
values = {field.name: getattr(config, field.name) for field in fields(config)}
|
|
36
|
+
elif isinstance(config, Mapping):
|
|
37
|
+
values = dict(config)
|
|
38
|
+
else:
|
|
39
|
+
raise TypeError("config must be an OpensplineConfig or mapping")
|
|
40
|
+
overrides = {
|
|
41
|
+
key: value
|
|
42
|
+
for key, value in {
|
|
43
|
+
"url": url,
|
|
44
|
+
"timeout": timeout,
|
|
45
|
+
"quality": quality,
|
|
46
|
+
"viewer_timeout": viewer_timeout,
|
|
47
|
+
}.items()
|
|
48
|
+
if value is not None
|
|
49
|
+
}
|
|
50
|
+
self.config = OpensplineConfig(**(values | overrides))
|
|
51
|
+
self.url = self.config.url
|
|
52
|
+
self.timeout = self.config.timeout
|
|
53
|
+
|
|
54
|
+
def avatar(self, portrait, quality: Quality | None = None):
|
|
55
|
+
quality = self.config.quality if quality is None else quality
|
|
56
|
+
if quality not in {"low", "high"}:
|
|
57
|
+
raise ValueError("quality must be low or high")
|
|
58
|
+
return AvatarSession(self, portrait, quality)
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class AvatarSession:
|
|
62
|
+
def __init__(self, client, portrait, quality):
|
|
63
|
+
self.client, self.portrait, self.quality = client, portrait, quality
|
|
64
|
+
self.id = None
|
|
65
|
+
self.viewer_url = None
|
|
66
|
+
self.session = None
|
|
67
|
+
self.ws = None
|
|
68
|
+
self.http = None
|
|
69
|
+
self.pending = {}
|
|
70
|
+
self.callbacks = []
|
|
71
|
+
self.epoch = 0
|
|
72
|
+
self.viewer_ready = asyncio.Event()
|
|
73
|
+
self.reader = None
|
|
74
|
+
self.heartbeat = None
|
|
75
|
+
self._dirty = False
|
|
76
|
+
self._closed = False
|
|
77
|
+
self._send_lock = asyncio.Lock()
|
|
78
|
+
self._media = None
|
|
79
|
+
|
|
80
|
+
async def __aenter__(self):
|
|
81
|
+
self.http = httpx.AsyncClient(base_url=self.client.url, timeout=self.client.timeout)
|
|
82
|
+
try:
|
|
83
|
+
if isinstance(self.portrait, (str, Path)):
|
|
84
|
+
raw = Path(self.portrait).read_bytes()
|
|
85
|
+
else:
|
|
86
|
+
raw = bytes(self.portrait)
|
|
87
|
+
response = await self.http.post(
|
|
88
|
+
"/v1/sessions",
|
|
89
|
+
files={"portrait": ("portrait", raw)},
|
|
90
|
+
data={"quality": self.quality},
|
|
91
|
+
)
|
|
92
|
+
self._check(response)
|
|
93
|
+
body = response.json()
|
|
94
|
+
self.id = body["id"]
|
|
95
|
+
self.token = body["publisher_token"]
|
|
96
|
+
self.session = {k: v for k, v in body.items() if k != "publisher_token"}
|
|
97
|
+
self.viewer_url = body["viewer_url"]
|
|
98
|
+
url = (
|
|
99
|
+
self.client.url.replace("https://", "wss://").replace("http://", "ws://")
|
|
100
|
+
+ f"/v1/sessions/{self.id}/audio"
|
|
101
|
+
)
|
|
102
|
+
self.ws = await connect(url, max_size=1024 * 1024, open_timeout=30)
|
|
103
|
+
await self.ws.send(json.dumps({"token": self.token}))
|
|
104
|
+
self.reader = asyncio.create_task(self._receive())
|
|
105
|
+
self.heartbeat = asyncio.create_task(self._heartbeat())
|
|
106
|
+
return self
|
|
107
|
+
except BaseException:
|
|
108
|
+
await self.close()
|
|
109
|
+
raise
|
|
110
|
+
|
|
111
|
+
async def __aexit__(self, typ, exc, tb):
|
|
112
|
+
try:
|
|
113
|
+
if typ is None and self._dirty:
|
|
114
|
+
await self.end_turn()
|
|
115
|
+
finally:
|
|
116
|
+
await self.close()
|
|
117
|
+
|
|
118
|
+
@staticmethod
|
|
119
|
+
def _check(response):
|
|
120
|
+
if response.is_error:
|
|
121
|
+
try:
|
|
122
|
+
error = response.json().get("error", {})
|
|
123
|
+
message = error.get("message", response.text)
|
|
124
|
+
code = error.get("code", "request")
|
|
125
|
+
except (ValueError, AttributeError):
|
|
126
|
+
message = response.text
|
|
127
|
+
code = "request"
|
|
128
|
+
raise error_from(code, message)
|
|
129
|
+
|
|
130
|
+
def on(self, callback):
|
|
131
|
+
self.callbacks.append(callback)
|
|
132
|
+
return lambda: self.callbacks.remove(callback)
|
|
133
|
+
|
|
134
|
+
async def _receive(self):
|
|
135
|
+
failure = ConnectionError("Session connection closed", "connection")
|
|
136
|
+
try:
|
|
137
|
+
async for raw in self.ws:
|
|
138
|
+
event = json.loads(raw)
|
|
139
|
+
self.epoch = max(self.epoch, event.get("epoch", self.epoch))
|
|
140
|
+
if event.get("type") == "viewer_ready" or event.get("viewer_ready"):
|
|
141
|
+
self.viewer_ready.set()
|
|
142
|
+
if event.get("type") == "viewer_disconnected":
|
|
143
|
+
self.viewer_ready.clear()
|
|
144
|
+
id = event.get("id")
|
|
145
|
+
future = self.pending.get(id)
|
|
146
|
+
if event.get("type") == "error":
|
|
147
|
+
error = error_from(
|
|
148
|
+
event.get("code", "error"), event.get("message", "Session error")
|
|
149
|
+
)
|
|
150
|
+
if future and not future.done():
|
|
151
|
+
future.set_exception(error)
|
|
152
|
+
elif not id:
|
|
153
|
+
failure = error
|
|
154
|
+
for f in self.pending.values():
|
|
155
|
+
if not f.done():
|
|
156
|
+
f.set_exception(error)
|
|
157
|
+
elif future and not future.done():
|
|
158
|
+
future.set_result(event)
|
|
159
|
+
for cb in list(self.callbacks):
|
|
160
|
+
try:
|
|
161
|
+
value = cb(event)
|
|
162
|
+
if inspect.isawaitable(value):
|
|
163
|
+
asyncio.create_task(value)
|
|
164
|
+
except Exception:
|
|
165
|
+
import logging
|
|
166
|
+
|
|
167
|
+
logging.getLogger(__name__).exception("Session event callback failed")
|
|
168
|
+
except Exception as exc:
|
|
169
|
+
failure = ConnectionError(str(exc), "connection")
|
|
170
|
+
finally:
|
|
171
|
+
for future in self.pending.values():
|
|
172
|
+
if not future.done():
|
|
173
|
+
future.set_exception(failure)
|
|
174
|
+
|
|
175
|
+
async def _heartbeat(self):
|
|
176
|
+
try:
|
|
177
|
+
while True:
|
|
178
|
+
await asyncio.sleep(15)
|
|
179
|
+
await self._request("ping")
|
|
180
|
+
except (asyncio.CancelledError, Exception):
|
|
181
|
+
return
|
|
182
|
+
|
|
183
|
+
async def _request(self, type, **data):
|
|
184
|
+
if self._closed or not self.ws:
|
|
185
|
+
raise ConnectionError("Session is not connected", "connection")
|
|
186
|
+
id = uuid.uuid4().hex
|
|
187
|
+
future = asyncio.get_running_loop().create_future()
|
|
188
|
+
self.pending[id] = future
|
|
189
|
+
try:
|
|
190
|
+
async with self._send_lock:
|
|
191
|
+
await self.ws.send(json.dumps({"type": type, "id": id, **data}))
|
|
192
|
+
return await asyncio.wait_for(future, self.client.timeout)
|
|
193
|
+
finally:
|
|
194
|
+
self.pending.pop(id, None)
|
|
195
|
+
|
|
196
|
+
async def wait_for_viewer(self, timeout=None):
|
|
197
|
+
await asyncio.wait_for(
|
|
198
|
+
self.viewer_ready.wait(),
|
|
199
|
+
self.client.config.viewer_timeout if timeout is None else timeout,
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
async def send_audio(self, data, *, sample_rate=24000, channels=1, encoding="pcm_s16le"):
|
|
203
|
+
if (
|
|
204
|
+
sample_rate not in {8000, 16000, 22050, 24000, 32000, 44100, 48000}
|
|
205
|
+
or channels not in {1, 2}
|
|
206
|
+
or encoding not in {"pcm_s16le", "pcm_f32le", "mp3"}
|
|
207
|
+
):
|
|
208
|
+
raise ValueError("Invalid audio format")
|
|
209
|
+
raw = bytes(data)
|
|
210
|
+
size = (
|
|
211
|
+
16384
|
|
212
|
+
if encoding == "mp3"
|
|
213
|
+
else sample_rate * channels * (4 if encoding == "pcm_f32le" else 2) // 10
|
|
214
|
+
)
|
|
215
|
+
epoch = self.epoch
|
|
216
|
+
for i in range(0, len(raw), size):
|
|
217
|
+
if epoch != self.epoch:
|
|
218
|
+
return
|
|
219
|
+
await self._request(
|
|
220
|
+
"audio",
|
|
221
|
+
epoch=epoch,
|
|
222
|
+
data=base64.b64encode(raw[i : i + size]).decode(),
|
|
223
|
+
format={"sample_rate": sample_rate, "channels": channels, "encoding": encoding},
|
|
224
|
+
)
|
|
225
|
+
self._dirty = True
|
|
226
|
+
|
|
227
|
+
async def stream(self, source, **format):
|
|
228
|
+
await self.wait_for_viewer()
|
|
229
|
+
try:
|
|
230
|
+
async for packet in source:
|
|
231
|
+
await self.send_audio(packet, **format)
|
|
232
|
+
await self.end_turn()
|
|
233
|
+
except BaseException:
|
|
234
|
+
with contextlib.suppress(Exception):
|
|
235
|
+
await self.interrupt()
|
|
236
|
+
raise
|
|
237
|
+
|
|
238
|
+
async def end_turn(self, drain=True):
|
|
239
|
+
result = await self._request("end_turn", drain=drain)
|
|
240
|
+
self._dirty = False
|
|
241
|
+
return result
|
|
242
|
+
|
|
243
|
+
async def interrupt(self):
|
|
244
|
+
# A separate request can overtake a backpressured audio WebSocket.
|
|
245
|
+
response = await self.http.post(
|
|
246
|
+
f"/v1/sessions/{self.id}/interrupt", headers=self._headers()
|
|
247
|
+
)
|
|
248
|
+
self._check(response)
|
|
249
|
+
self.epoch = response.json()["epoch"]
|
|
250
|
+
self._dirty = False
|
|
251
|
+
|
|
252
|
+
def _headers(self):
|
|
253
|
+
return {"Authorization": "Bearer " + self.token}
|
|
254
|
+
|
|
255
|
+
async def viewer(self):
|
|
256
|
+
response = await self.http.post(f"/v1/sessions/{self.id}/viewer", headers=self._headers())
|
|
257
|
+
self._check(response)
|
|
258
|
+
self.session = response.json()
|
|
259
|
+
self.viewer_url = self.session["viewer_url"]
|
|
260
|
+
return self.session
|
|
261
|
+
|
|
262
|
+
async def play_file(self, path):
|
|
263
|
+
await self.wait_for_viewer()
|
|
264
|
+
with Path(path).open("rb") as f:
|
|
265
|
+
response = await self.http.post(
|
|
266
|
+
f"/v1/sessions/{self.id}/file",
|
|
267
|
+
headers=self._headers(),
|
|
268
|
+
files={"audio": (Path(path).name, f)},
|
|
269
|
+
)
|
|
270
|
+
self._check(response)
|
|
271
|
+
self._dirty = True
|
|
272
|
+
await self.end_turn()
|
|
273
|
+
|
|
274
|
+
async def media(self):
|
|
275
|
+
from .media import MediaReceiver
|
|
276
|
+
|
|
277
|
+
if self._media is None:
|
|
278
|
+
self._media = MediaReceiver(self)
|
|
279
|
+
await self._media.start()
|
|
280
|
+
return self._media
|
|
281
|
+
|
|
282
|
+
async def close(self):
|
|
283
|
+
if self._closed:
|
|
284
|
+
return
|
|
285
|
+
self._closed = True
|
|
286
|
+
if self._media:
|
|
287
|
+
await self._media.close()
|
|
288
|
+
if self.heartbeat:
|
|
289
|
+
self.heartbeat.cancel()
|
|
290
|
+
# Explicit delete is idempotent and releases the worker even if WebSocket setup failed.
|
|
291
|
+
if self.id and self.http:
|
|
292
|
+
with contextlib.suppress(Exception):
|
|
293
|
+
await self.http.delete(f"/v1/sessions/{self.id}", headers=self._headers())
|
|
294
|
+
if self.ws:
|
|
295
|
+
await self.ws.close()
|
|
296
|
+
if self.reader:
|
|
297
|
+
self.reader.cancel()
|
|
298
|
+
with contextlib.suppress(asyncio.CancelledError):
|
|
299
|
+
await self.reader
|
|
300
|
+
if self.http:
|
|
301
|
+
await self.http.aclose()
|
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
"""Validated, dependency-free client configuration."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
import os
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from typing import Literal
|
|
7
|
+
from urllib.parse import urlsplit
|
|
8
|
+
|
|
9
|
+
Quality = Literal["low", "high"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@dataclass(frozen=True)
|
|
13
|
+
class OpensplineConfig:
|
|
14
|
+
url: str = field(default_factory=lambda: os.getenv("OPENSPLINE_URL", "http://localhost:7860"))
|
|
15
|
+
quality: Quality = "low"
|
|
16
|
+
timeout: float = 180
|
|
17
|
+
viewer_timeout: float = 60
|
|
18
|
+
|
|
19
|
+
def __post_init__(self):
|
|
20
|
+
if not isinstance(self.url, str):
|
|
21
|
+
raise ValueError("url must be an HTTP or HTTPS URL")
|
|
22
|
+
try:
|
|
23
|
+
parsed = urlsplit(self.url)
|
|
24
|
+
valid = parsed.scheme in {"http", "https"} and parsed.hostname and parsed.port != 0
|
|
25
|
+
except ValueError:
|
|
26
|
+
valid = False
|
|
27
|
+
if (
|
|
28
|
+
not valid
|
|
29
|
+
or parsed.username is not None
|
|
30
|
+
or parsed.password is not None
|
|
31
|
+
or parsed.query
|
|
32
|
+
or parsed.fragment
|
|
33
|
+
or any(c.isspace() for c in self.url)
|
|
34
|
+
):
|
|
35
|
+
raise ValueError(
|
|
36
|
+
"url must be an HTTP or HTTPS URL without credentials, query, or fragment"
|
|
37
|
+
)
|
|
38
|
+
if self.quality not in ("low", "high"):
|
|
39
|
+
raise ValueError("quality must be low or high")
|
|
40
|
+
for name in ("timeout", "viewer_timeout"):
|
|
41
|
+
value = getattr(self, name)
|
|
42
|
+
if (
|
|
43
|
+
isinstance(value, bool)
|
|
44
|
+
or not isinstance(value, (int, float))
|
|
45
|
+
or not math.isfinite(value)
|
|
46
|
+
or value <= 0
|
|
47
|
+
):
|
|
48
|
+
raise ValueError(f"{name} must be a finite positive number of seconds")
|
|
49
|
+
object.__setattr__(self, "url", self.url.rstrip("/"))
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
class OpensplineError(Exception):
|
|
2
|
+
def __init__(self, message, code="error"):
|
|
3
|
+
super().__init__(message)
|
|
4
|
+
self.code = code
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class CapacityError(OpensplineError):
|
|
8
|
+
pass
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class ConfigurationError(OpensplineError):
|
|
12
|
+
pass
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class ConnectionError(OpensplineError):
|
|
16
|
+
pass
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class InferenceError(OpensplineError):
|
|
20
|
+
pass
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def error_from(code, message):
|
|
24
|
+
cls = {
|
|
25
|
+
"capacity": CapacityError,
|
|
26
|
+
"configuration": ConfigurationError,
|
|
27
|
+
"connection": ConnectionError,
|
|
28
|
+
"inference": InferenceError,
|
|
29
|
+
}.get(code, OpensplineError)
|
|
30
|
+
return cls(message, code)
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
"""Optional integrations. Import only the adapter your application uses."""
|
|
2
|
+
|
|
3
|
+
from .elevenlabs import ElevenLabsAgents
|
|
4
|
+
from .events import GeminiLive, GoogleADK, OpenAIAgents, OpenAIRealtime
|
|
5
|
+
|
|
6
|
+
__all__ = ["ElevenLabsAgents", "GeminiLive", "GoogleADK", "OpenAIAgents", "OpenAIRealtime"]
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
"""ElevenLabs Agents WebSocket events, without a provider SDK dependency."""
|
|
2
|
+
|
|
3
|
+
import base64
|
|
4
|
+
import struct
|
|
5
|
+
|
|
6
|
+
from .events import EventAdapter, value
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def audio_rate(audio_format):
|
|
10
|
+
if audio_format == "ulaw_8000":
|
|
11
|
+
return 8000
|
|
12
|
+
if audio_format not in {f"pcm_{rate}" for rate in (8000, 16000, 22050, 24000, 44100, 48000)}:
|
|
13
|
+
raise ValueError("Unsupported ElevenLabs audio format; use PCM or ulaw_8000.")
|
|
14
|
+
return int(audio_format[4:])
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def decode_audio(encoded, audio_format):
|
|
18
|
+
raw = base64.b64decode(encoded, validate=True)
|
|
19
|
+
if audio_format == "ulaw_8000":
|
|
20
|
+
samples = []
|
|
21
|
+
for byte in raw:
|
|
22
|
+
byte = ~byte & 255
|
|
23
|
+
sample = (((byte & 15) << 3) + 132) << ((byte >> 4) & 7)
|
|
24
|
+
samples.append(132 - sample if byte & 128 else sample - 132)
|
|
25
|
+
return struct.pack(f"<{len(samples)}h", *samples)
|
|
26
|
+
return raw
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class ElevenLabsAgents(EventAdapter):
|
|
30
|
+
"""Forward decoded WebSocket events; the application owns pings, tools, and input.
|
|
31
|
+
|
|
32
|
+
Include the initiation metadata event to select the negotiated audio format.
|
|
33
|
+
If attaching after initiation, pass that format explicitly.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
def __init__(self, avatar, *, audio_format="pcm_16000"):
|
|
37
|
+
super().__init__(avatar)
|
|
38
|
+
self.audio_format = audio_format
|
|
39
|
+
self.sample_rate = audio_rate(audio_format)
|
|
40
|
+
self.interrupted_id = -1
|
|
41
|
+
self.pending = False
|
|
42
|
+
|
|
43
|
+
async def handle(self, event):
|
|
44
|
+
kind = value(event, "type")
|
|
45
|
+
if kind == "conversation_initiation_metadata":
|
|
46
|
+
metadata = value(event, "conversation_initiation_metadata_event")
|
|
47
|
+
self.audio_format = value(metadata, "agent_output_audio_format")
|
|
48
|
+
self.sample_rate = audio_rate(self.audio_format)
|
|
49
|
+
elif kind == "interruption":
|
|
50
|
+
interruption = value(event, "interruption_event")
|
|
51
|
+
self.interrupted_id = max(self.interrupted_id, value(interruption, "event_id", -1))
|
|
52
|
+
self.pending = False
|
|
53
|
+
await self.avatar.interrupt()
|
|
54
|
+
elif kind == "audio":
|
|
55
|
+
audio = value(event, "audio_event")
|
|
56
|
+
if value(audio, "event_id", 0) <= self.interrupted_id:
|
|
57
|
+
return
|
|
58
|
+
raw = decode_audio(value(audio, "audio_base_64"), self.audio_format)
|
|
59
|
+
if raw:
|
|
60
|
+
await self.avatar.send_audio(raw, sample_rate=self.sample_rate)
|
|
61
|
+
self.pending = True
|
|
62
|
+
if value(audio, "is_final", False):
|
|
63
|
+
await self._finish()
|
|
64
|
+
elif kind == "agent_response_complete":
|
|
65
|
+
complete = value(event, "agent_response_complete_event")
|
|
66
|
+
if value(complete, "event_id", 0) > self.interrupted_id:
|
|
67
|
+
await self._finish()
|
|
68
|
+
|
|
69
|
+
async def _finish(self):
|
|
70
|
+
if self.pending:
|
|
71
|
+
await self.avatar.end_turn(drain=False)
|
|
72
|
+
self.pending = False
|
|
@@ -0,0 +1,153 @@
|
|
|
1
|
+
import base64
|
|
2
|
+
import re
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def value(obj, key, default=None):
|
|
6
|
+
return obj.get(key, default) if isinstance(obj, dict) else getattr(obj, key, default)
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class EventAdapter:
|
|
10
|
+
def __init__(self, avatar):
|
|
11
|
+
self.avatar = avatar
|
|
12
|
+
|
|
13
|
+
async def wrap(self, events):
|
|
14
|
+
"""Preserve every provider event for the application's existing event loop."""
|
|
15
|
+
async for event in events:
|
|
16
|
+
await self.handle(event)
|
|
17
|
+
yield event
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class OpenAIRealtime(EventAdapter):
|
|
21
|
+
def __init__(self, avatar, connection=None):
|
|
22
|
+
super().__init__(avatar)
|
|
23
|
+
self.connection = connection
|
|
24
|
+
self.item = None
|
|
25
|
+
self.content_index = 0
|
|
26
|
+
self.played = 0
|
|
27
|
+
self.generated = 0
|
|
28
|
+
self.base = 0
|
|
29
|
+
self.unsubscribe = avatar.on(self._playback)
|
|
30
|
+
|
|
31
|
+
def _playback(self, event):
|
|
32
|
+
if event.get("type") == "playback":
|
|
33
|
+
self.played = event["samples"] / 48000
|
|
34
|
+
|
|
35
|
+
async def handle(self, event):
|
|
36
|
+
kind = value(event, "type")
|
|
37
|
+
if kind == "response.output_audio.delta":
|
|
38
|
+
item = value(event, "item_id")
|
|
39
|
+
if item != self.item:
|
|
40
|
+
self.item = item
|
|
41
|
+
self.base = self.generated
|
|
42
|
+
self.content_index = value(event, "content_index", 0)
|
|
43
|
+
raw = base64.b64decode(value(event, "delta"))
|
|
44
|
+
self.generated += len(raw) / 48000
|
|
45
|
+
await self.avatar.send_audio(raw, sample_rate=24000)
|
|
46
|
+
elif kind == "response.output_audio.done":
|
|
47
|
+
await self.avatar.end_turn(drain=False)
|
|
48
|
+
elif kind == "input_audio_buffer.speech_started":
|
|
49
|
+
heard = max(0, self.played - self.base)
|
|
50
|
+
await self.avatar.interrupt()
|
|
51
|
+
if self.connection and self.item:
|
|
52
|
+
await self.connection.send(
|
|
53
|
+
{
|
|
54
|
+
"type": "conversation.item.truncate",
|
|
55
|
+
"item_id": self.item,
|
|
56
|
+
"content_index": self.content_index,
|
|
57
|
+
"audio_end_ms": int(heard * 1000),
|
|
58
|
+
}
|
|
59
|
+
)
|
|
60
|
+
self.item = None
|
|
61
|
+
self.generated = self.played = self.base = 0
|
|
62
|
+
|
|
63
|
+
def close(self):
|
|
64
|
+
self.unsubscribe()
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class OpenAIAgents(EventAdapter):
|
|
68
|
+
"""Use .tracker in RealtimeRunner.run(model_config={"playback_tracker": ...})."""
|
|
69
|
+
|
|
70
|
+
def __init__(self, avatar, playback_tracker=None):
|
|
71
|
+
from agents.realtime import RealtimePlaybackTracker
|
|
72
|
+
|
|
73
|
+
super().__init__(avatar)
|
|
74
|
+
self.tracker = playback_tracker or RealtimePlaybackTracker()
|
|
75
|
+
self.spans = []
|
|
76
|
+
self.generated = 0
|
|
77
|
+
self.played = 0
|
|
78
|
+
self.unsubscribe = avatar.on(self._playback)
|
|
79
|
+
|
|
80
|
+
def _playback(self, event):
|
|
81
|
+
if event.get("type") != "playback":
|
|
82
|
+
return
|
|
83
|
+
new = event["samples"] / 48
|
|
84
|
+
for start, end, item, index in self.spans:
|
|
85
|
+
duration = max(0, min(new, end) - max(self.played, start))
|
|
86
|
+
if duration:
|
|
87
|
+
self.tracker.on_play_ms(item, index, duration)
|
|
88
|
+
self.played = max(self.played, new)
|
|
89
|
+
self.spans = [span for span in self.spans if span[1] > self.played]
|
|
90
|
+
|
|
91
|
+
async def handle(self, event):
|
|
92
|
+
kind = value(event, "type")
|
|
93
|
+
if kind == "audio":
|
|
94
|
+
audio = value(event, "audio")
|
|
95
|
+
raw = value(audio, "data", audio)
|
|
96
|
+
duration = len(raw) / 48
|
|
97
|
+
self.spans.append(
|
|
98
|
+
(
|
|
99
|
+
self.generated,
|
|
100
|
+
self.generated + duration,
|
|
101
|
+
value(event, "item_id"),
|
|
102
|
+
value(event, "content_index", 0),
|
|
103
|
+
)
|
|
104
|
+
)
|
|
105
|
+
self.generated += duration
|
|
106
|
+
await self.avatar.send_audio(raw, sample_rate=24000)
|
|
107
|
+
elif kind == "audio_end":
|
|
108
|
+
await self.avatar.end_turn(drain=False)
|
|
109
|
+
elif kind == "audio_interrupted":
|
|
110
|
+
await self.avatar.interrupt()
|
|
111
|
+
self.spans = []
|
|
112
|
+
self.generated = self.played = 0
|
|
113
|
+
|
|
114
|
+
def close(self):
|
|
115
|
+
self.unsubscribe()
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
class GeminiLive(EventAdapter):
|
|
119
|
+
async def handle(self, event):
|
|
120
|
+
content = value(event, "server_content", value(event, "serverContent"))
|
|
121
|
+
if not content:
|
|
122
|
+
return
|
|
123
|
+
if value(content, "interrupted", False):
|
|
124
|
+
await self.avatar.interrupt()
|
|
125
|
+
return
|
|
126
|
+
turn = value(content, "model_turn", value(content, "modelTurn"))
|
|
127
|
+
for part in value(turn, "parts", []) or []:
|
|
128
|
+
blob = value(part, "inline_data", value(part, "inlineData"))
|
|
129
|
+
mime = value(blob, "mime_type", value(blob, "mimeType", ""))
|
|
130
|
+
if blob and mime.startswith("audio/"):
|
|
131
|
+
raw = value(blob, "data")
|
|
132
|
+
raw = base64.b64decode(raw) if isinstance(raw, str) else raw
|
|
133
|
+
match = re.search(r"rate=(\d+)", mime)
|
|
134
|
+
await self.avatar.send_audio(raw, sample_rate=int(match[1]) if match else 24000)
|
|
135
|
+
if value(content, "turn_complete", value(content, "turnComplete", False)):
|
|
136
|
+
await self.avatar.end_turn(drain=False)
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
class GoogleADK(EventAdapter):
|
|
140
|
+
async def handle(self, event):
|
|
141
|
+
if value(event, "interrupted", False):
|
|
142
|
+
await self.avatar.interrupt()
|
|
143
|
+
return
|
|
144
|
+
for part in value(value(event, "content"), "parts", []) or []:
|
|
145
|
+
blob = value(part, "inline_data")
|
|
146
|
+
mime = value(blob, "mime_type", "")
|
|
147
|
+
if blob and mime.startswith("audio/"):
|
|
148
|
+
match = re.search(r"rate=(\d+)", mime)
|
|
149
|
+
await self.avatar.send_audio(
|
|
150
|
+
value(blob, "data"), sample_rate=int(match[1]) if match else 24000
|
|
151
|
+
)
|
|
152
|
+
if value(event, "turn_complete", False):
|
|
153
|
+
await self.avatar.end_turn(drain=False)
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
"""LiveKit native avatar plugin. Install the server's LiveKit extra as well."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import secrets
|
|
5
|
+
|
|
6
|
+
from livekit import api
|
|
7
|
+
from livekit.agents.voice.avatar import AvatarSession as BaseAvatarSession
|
|
8
|
+
from livekit.agents.voice.avatar import DataStreamAudioOutput
|
|
9
|
+
|
|
10
|
+
from openspline import Openspline
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class AvatarSession(BaseAvatarSession):
|
|
14
|
+
def __init__(
|
|
15
|
+
self, portrait, quality=None, client=None, url=None, api_key=None, api_secret=None
|
|
16
|
+
):
|
|
17
|
+
super().__init__()
|
|
18
|
+
self.client = client or Openspline()
|
|
19
|
+
self.portrait = portrait
|
|
20
|
+
self.quality = quality
|
|
21
|
+
self.url = url or os.environ.get("LIVEKIT_URL")
|
|
22
|
+
self.key = api_key or os.environ.get("LIVEKIT_API_KEY")
|
|
23
|
+
self.secret = api_secret or os.environ.get("LIVEKIT_API_SECRET")
|
|
24
|
+
self.identity = "openspline-" + secrets.token_hex(6)
|
|
25
|
+
self.avatar = None
|
|
26
|
+
|
|
27
|
+
@property
|
|
28
|
+
def avatar_identity(self):
|
|
29
|
+
return self.identity
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
def provider(self):
|
|
33
|
+
return "openspline"
|
|
34
|
+
|
|
35
|
+
async def start(self, agent_session, room):
|
|
36
|
+
await super().start(agent_session, room)
|
|
37
|
+
self.avatar = await self.client.avatar(self.portrait, self.quality).__aenter__()
|
|
38
|
+
token = (
|
|
39
|
+
api.AccessToken(self.key, self.secret)
|
|
40
|
+
.with_identity(self.identity)
|
|
41
|
+
.with_kind("agent")
|
|
42
|
+
.with_attributes({"lk.publish_on_behalf": room.local_participant.identity})
|
|
43
|
+
.with_grants(api.VideoGrants(room_join=True, room=room.name))
|
|
44
|
+
.to_jwt()
|
|
45
|
+
)
|
|
46
|
+
try:
|
|
47
|
+
response = await self.avatar.http.post(
|
|
48
|
+
f"/v1/sessions/{self.avatar.id}/livekit",
|
|
49
|
+
headers=self.avatar._headers(),
|
|
50
|
+
json={
|
|
51
|
+
"url": self.url,
|
|
52
|
+
"token": token,
|
|
53
|
+
"sender_identity": room.local_participant.identity,
|
|
54
|
+
},
|
|
55
|
+
)
|
|
56
|
+
self.avatar._check(response)
|
|
57
|
+
agent_session.output.audio = DataStreamAudioOutput(
|
|
58
|
+
room, destination_identity=self.identity, wait_playback_start=True
|
|
59
|
+
)
|
|
60
|
+
except BaseException:
|
|
61
|
+
await self.aclose()
|
|
62
|
+
raise
|
|
63
|
+
|
|
64
|
+
async def aclose(self):
|
|
65
|
+
if self.avatar:
|
|
66
|
+
await self.avatar.close()
|
|
67
|
+
self.avatar = None
|
|
68
|
+
await super().aclose()
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
"""Place after TTS/realtime output and before transport.output()."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
|
|
5
|
+
from pipecat.frames.frames import (
|
|
6
|
+
CancelFrame,
|
|
7
|
+
EndFrame,
|
|
8
|
+
Frame,
|
|
9
|
+
InterruptionFrame,
|
|
10
|
+
OutputImageRawFrame,
|
|
11
|
+
StartFrame,
|
|
12
|
+
TTSAudioRawFrame,
|
|
13
|
+
TTSStoppedFrame,
|
|
14
|
+
)
|
|
15
|
+
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
|
16
|
+
|
|
17
|
+
from openspline import Openspline
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class AvatarProcessor(FrameProcessor):
|
|
21
|
+
def __init__(self, portrait, quality=None, client=None, **kwargs):
|
|
22
|
+
super().__init__(**kwargs)
|
|
23
|
+
self.client = client or Openspline()
|
|
24
|
+
self.portrait = portrait
|
|
25
|
+
self.quality = quality
|
|
26
|
+
self.avatar = None
|
|
27
|
+
self.task = None
|
|
28
|
+
|
|
29
|
+
async def _output(self):
|
|
30
|
+
import av
|
|
31
|
+
|
|
32
|
+
resampler = av.AudioResampler(format="s16", layout="mono", rate=48000)
|
|
33
|
+
media = await self.avatar.media()
|
|
34
|
+
async for kind, frame in media:
|
|
35
|
+
if kind == "audio":
|
|
36
|
+
for out in resampler.resample(frame):
|
|
37
|
+
await self.push_frame(
|
|
38
|
+
TTSAudioRawFrame(
|
|
39
|
+
audio=out.to_ndarray().tobytes(), sample_rate=48000, num_channels=1
|
|
40
|
+
)
|
|
41
|
+
)
|
|
42
|
+
media.acknowledge()
|
|
43
|
+
else:
|
|
44
|
+
out = OutputImageRawFrame(
|
|
45
|
+
image=frame.to_ndarray(format="rgb24").tobytes(),
|
|
46
|
+
size=(frame.width, frame.height),
|
|
47
|
+
format="RGB",
|
|
48
|
+
)
|
|
49
|
+
if frame.pts is not None:
|
|
50
|
+
out.pts = int(frame.pts * frame.time_base * 1e9)
|
|
51
|
+
await self.push_frame(out)
|
|
52
|
+
|
|
53
|
+
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
54
|
+
await super().process_frame(frame, direction)
|
|
55
|
+
if direction != FrameDirection.DOWNSTREAM:
|
|
56
|
+
await self.push_frame(frame, direction)
|
|
57
|
+
return
|
|
58
|
+
if isinstance(frame, StartFrame):
|
|
59
|
+
self.avatar = await self.client.avatar(self.portrait, self.quality).__aenter__()
|
|
60
|
+
await self.push_frame(frame, direction)
|
|
61
|
+
self.task = asyncio.create_task(self._output())
|
|
62
|
+
return
|
|
63
|
+
if isinstance(frame, TTSAudioRawFrame):
|
|
64
|
+
await self.avatar.send_audio(
|
|
65
|
+
frame.audio, sample_rate=frame.sample_rate, channels=frame.num_channels
|
|
66
|
+
)
|
|
67
|
+
return
|
|
68
|
+
if isinstance(frame, TTSStoppedFrame):
|
|
69
|
+
await self.avatar.end_turn()
|
|
70
|
+
if isinstance(frame, InterruptionFrame):
|
|
71
|
+
await self.avatar.interrupt()
|
|
72
|
+
if isinstance(frame, (EndFrame, CancelFrame)):
|
|
73
|
+
await self._close()
|
|
74
|
+
await self.push_frame(frame, direction)
|
|
75
|
+
|
|
76
|
+
async def _close(self):
|
|
77
|
+
if self.task:
|
|
78
|
+
self.task.cancel()
|
|
79
|
+
await asyncio.gather(self.task, return_exceptions=True)
|
|
80
|
+
self.task = None
|
|
81
|
+
if self.avatar:
|
|
82
|
+
await self.avatar.close()
|
|
83
|
+
self.avatar = None
|
|
84
|
+
|
|
85
|
+
async def cleanup(self):
|
|
86
|
+
await self._close()
|
|
87
|
+
await super().cleanup()
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""TEN extension using its public AsyncExtension and AudioFrame interfaces."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
|
|
5
|
+
from ten_runtime import AsyncExtension, CmdResult, Data, StatusCode
|
|
6
|
+
|
|
7
|
+
from openspline import Openspline
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class AvatarExtension(AsyncExtension):
|
|
11
|
+
async def on_start(self, env):
|
|
12
|
+
raw, error = await env.get_property_to_json("")
|
|
13
|
+
if error:
|
|
14
|
+
raise ValueError("Cannot read openspline properties")
|
|
15
|
+
config = json.loads(raw)
|
|
16
|
+
self.avatar = (
|
|
17
|
+
await Openspline().avatar(config["portrait"], config.get("quality", "low")).__aenter__()
|
|
18
|
+
)
|
|
19
|
+
descriptor = Data.create("avatar_session")
|
|
20
|
+
descriptor.set_property_from_json("session", json.dumps(self.avatar.session))
|
|
21
|
+
await env.send_data(descriptor)
|
|
22
|
+
|
|
23
|
+
async def on_audio_frame(self, env, frame):
|
|
24
|
+
await self.avatar.send_audio(
|
|
25
|
+
bytes(frame.get_buf()),
|
|
26
|
+
sample_rate=frame.get_sample_rate(),
|
|
27
|
+
channels=frame.get_number_of_channels(),
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
async def on_data(self, env, data):
|
|
31
|
+
if data.get_name() == "tts_audio_end":
|
|
32
|
+
await self.avatar.end_turn(drain=False)
|
|
33
|
+
elif data.get_name() == "flush":
|
|
34
|
+
await self.avatar.interrupt()
|
|
35
|
+
|
|
36
|
+
async def on_cmd(self, env, cmd):
|
|
37
|
+
if cmd.get_name() == "flush":
|
|
38
|
+
await self.avatar.interrupt()
|
|
39
|
+
elif cmd.get_name() == "end_turn":
|
|
40
|
+
await self.avatar.end_turn()
|
|
41
|
+
await env.return_result(CmdResult.create(StatusCode.OK, cmd))
|
|
42
|
+
|
|
43
|
+
async def on_stop(self, env):
|
|
44
|
+
if getattr(self, "avatar", None):
|
|
45
|
+
await self.avatar.close()
|
|
@@ -0,0 +1,83 @@
|
|
|
1
|
+
"""Vocode output device for the pinned current AudioChunk worker API."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
from collections import deque
|
|
5
|
+
|
|
6
|
+
from vocode.streaming.models.audio import AudioEncoding
|
|
7
|
+
from vocode.streaming.output_device.abstract_output_device import AbstractOutputDevice
|
|
8
|
+
from vocode.streaming.output_device.audio_chunk import ChunkState
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class AvatarOutputDevice(AbstractOutputDevice):
|
|
12
|
+
"""Use inside an open avatar context; its viewer owns output playback."""
|
|
13
|
+
|
|
14
|
+
def __init__(self, avatar, sampling_rate=24000, max_pending_seconds=10):
|
|
15
|
+
super().__init__(sampling_rate, AudioEncoding.LINEAR16)
|
|
16
|
+
self._input_queue = asyncio.Queue(maxsize=64)
|
|
17
|
+
self.avatar = avatar
|
|
18
|
+
self.pending = deque()
|
|
19
|
+
self.samples = 0
|
|
20
|
+
self.played = 0
|
|
21
|
+
self.limit = int(max_pending_seconds * 48000)
|
|
22
|
+
self.changed = asyncio.Event()
|
|
23
|
+
self.unsubscribe = avatar.on(self._feedback)
|
|
24
|
+
self.tasks = set()
|
|
25
|
+
|
|
26
|
+
def _feedback(self, event):
|
|
27
|
+
if event.get("type") == "playback":
|
|
28
|
+
self.played = event["samples"]
|
|
29
|
+
while self.pending and self.pending[0][0] <= self.played:
|
|
30
|
+
_, event = self.pending.popleft()
|
|
31
|
+
chunk = event.payload
|
|
32
|
+
if event.is_interrupted():
|
|
33
|
+
chunk.state = ChunkState.INTERRUPTED
|
|
34
|
+
chunk.on_interrupt()
|
|
35
|
+
else:
|
|
36
|
+
chunk.state = ChunkState.PLAYED
|
|
37
|
+
chunk.on_play()
|
|
38
|
+
self.changed.set()
|
|
39
|
+
elif event.get("type") == "interrupted":
|
|
40
|
+
self._clear()
|
|
41
|
+
|
|
42
|
+
def _clear(self):
|
|
43
|
+
while self.pending:
|
|
44
|
+
_, event = self.pending.popleft()
|
|
45
|
+
event.payload.state = ChunkState.INTERRUPTED
|
|
46
|
+
event.payload.on_interrupt()
|
|
47
|
+
self.samples = self.played = 0
|
|
48
|
+
self.changed.set()
|
|
49
|
+
|
|
50
|
+
async def _run_loop(self):
|
|
51
|
+
while True:
|
|
52
|
+
event = await self._input_queue.get()
|
|
53
|
+
if event.is_interrupted():
|
|
54
|
+
event.payload.state = ChunkState.INTERRUPTED
|
|
55
|
+
event.payload.on_interrupt()
|
|
56
|
+
continue
|
|
57
|
+
while self.samples - self.played > self.limit:
|
|
58
|
+
self.changed.clear()
|
|
59
|
+
await self.changed.wait()
|
|
60
|
+
self.samples += round(len(event.payload.data) / 2 / self.sampling_rate * 48000)
|
|
61
|
+
self.pending.append((self.samples, event))
|
|
62
|
+
await self.avatar.send_audio(event.payload.data, sample_rate=self.sampling_rate)
|
|
63
|
+
# Vocode output chunks have no turn boundary. Flush each chunk, retaining
|
|
64
|
+
# model history. This favors correctness for short utterances over throughput.
|
|
65
|
+
await self.avatar.end_turn(drain=False)
|
|
66
|
+
|
|
67
|
+
def interrupt(self):
|
|
68
|
+
self._clear()
|
|
69
|
+
while not self._input_queue.empty():
|
|
70
|
+
event = self._input_queue.get_nowait()
|
|
71
|
+
event.payload.state = ChunkState.INTERRUPTED
|
|
72
|
+
event.payload.on_interrupt()
|
|
73
|
+
task = asyncio.create_task(self.avatar.interrupt())
|
|
74
|
+
self.tasks.add(task)
|
|
75
|
+
task.add_done_callback(self.tasks.discard)
|
|
76
|
+
|
|
77
|
+
async def terminate(self):
|
|
78
|
+
self.unsubscribe()
|
|
79
|
+
self._clear()
|
|
80
|
+
for task in self.tasks:
|
|
81
|
+
task.cancel()
|
|
82
|
+
await asyncio.gather(*self.tasks, return_exceptions=True)
|
|
83
|
+
await super().terminate()
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
"""Optional native media receiver for framework output adapters."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import json
|
|
5
|
+
|
|
6
|
+
from aiortc import RTCConfiguration, RTCIceServer, RTCPeerConnection, RTCSessionDescription
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class MediaReceiver:
|
|
10
|
+
def __init__(self, avatar):
|
|
11
|
+
self.avatar = avatar
|
|
12
|
+
self.pc = RTCPeerConnection()
|
|
13
|
+
self.frames = asyncio.Queue(100)
|
|
14
|
+
self.tasks = set()
|
|
15
|
+
self.clock = None
|
|
16
|
+
|
|
17
|
+
async def start(self):
|
|
18
|
+
response = await self.avatar.http.get(
|
|
19
|
+
f"/v1/sessions/{self.avatar.id}/ice",
|
|
20
|
+
headers={"Authorization": "Bearer " + self.avatar.session["token"]},
|
|
21
|
+
)
|
|
22
|
+
self.avatar._check(response)
|
|
23
|
+
await self.pc.close()
|
|
24
|
+
self.pc = RTCPeerConnection(
|
|
25
|
+
RTCConfiguration(iceServers=[RTCIceServer(**s) for s in response.json()["iceServers"]])
|
|
26
|
+
)
|
|
27
|
+
self.pc.addTransceiver("audio", direction="recvonly")
|
|
28
|
+
self.pc.addTransceiver("video", direction="recvonly")
|
|
29
|
+
self.channel = self.pc.createDataChannel("openspline")
|
|
30
|
+
|
|
31
|
+
@self.channel.on("open")
|
|
32
|
+
def opened():
|
|
33
|
+
self.channel.send(json.dumps({"type": "ready"}))
|
|
34
|
+
|
|
35
|
+
@self.channel.on("message")
|
|
36
|
+
def message(raw):
|
|
37
|
+
event = json.loads(raw)
|
|
38
|
+
if event["type"] == "clock":
|
|
39
|
+
self.clock = event
|
|
40
|
+
if event["type"] == "interrupted":
|
|
41
|
+
self.clock = None
|
|
42
|
+
while not self.frames.empty():
|
|
43
|
+
self.frames.get_nowait()
|
|
44
|
+
|
|
45
|
+
@self.pc.on("track")
|
|
46
|
+
def track(track):
|
|
47
|
+
async def receive():
|
|
48
|
+
while True:
|
|
49
|
+
frame = await track.recv()
|
|
50
|
+
await self.frames.put((track.kind, frame))
|
|
51
|
+
|
|
52
|
+
task = asyncio.create_task(receive())
|
|
53
|
+
self.tasks.add(task)
|
|
54
|
+
task.add_done_callback(self.tasks.discard)
|
|
55
|
+
|
|
56
|
+
await self.pc.setLocalDescription(await self.pc.createOffer())
|
|
57
|
+
s = self.avatar.session
|
|
58
|
+
response = await self.avatar.http.post(
|
|
59
|
+
f"/v1/sessions/{self.avatar.id}/offer",
|
|
60
|
+
headers={"Authorization": "Bearer " + s["token"]},
|
|
61
|
+
json={"type": self.pc.localDescription.type, "sdp": self.pc.localDescription.sdp},
|
|
62
|
+
)
|
|
63
|
+
self.avatar._check(response)
|
|
64
|
+
await self.pc.setRemoteDescription(RTCSessionDescription(**response.json()))
|
|
65
|
+
|
|
66
|
+
def acknowledge(self):
|
|
67
|
+
if self.clock and self.channel.readyState == "open":
|
|
68
|
+
self.channel.send(
|
|
69
|
+
json.dumps(
|
|
70
|
+
{
|
|
71
|
+
"type": "played",
|
|
72
|
+
"epoch": self.clock["epoch"],
|
|
73
|
+
"samples": self.clock["samples"],
|
|
74
|
+
}
|
|
75
|
+
)
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
def __aiter__(self):
|
|
79
|
+
return self
|
|
80
|
+
|
|
81
|
+
async def __anext__(self):
|
|
82
|
+
return await self.frames.get()
|
|
83
|
+
|
|
84
|
+
async def close(self):
|
|
85
|
+
for task in self.tasks:
|
|
86
|
+
task.cancel()
|
|
87
|
+
await asyncio.gather(*self.tasks, return_exceptions=True)
|
|
88
|
+
await self.pc.close()
|
|
File without changes
|