py-app-runner 0.5.49.dev0__py3-none-any.whl
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.
- py_app_runner/__init__.py +11 -0
- py_app_runner/audit/__init__.py +29 -0
- py_app_runner/audit/_service.py +91 -0
- py_app_runner/audit/_service_args.py +44 -0
- py_app_runner/audit/audit.py +319 -0
- py_app_runner/audit/commands.py +151 -0
- py_app_runner/audit/diff.py +202 -0
- py_app_runner/audit/errors.py +8 -0
- py_app_runner/audit/event.py +130 -0
- py_app_runner/audit/store.py +134 -0
- py_app_runner/bridge/__init__.py +0 -0
- py_app_runner/bridge/_service.py +265 -0
- py_app_runner/bridge/_service_args.py +24 -0
- py_app_runner/bridge/api.py +138 -0
- py_app_runner/bridge/encoders/__init__.py +5 -0
- py_app_runner/bridge/encoders/base.py +24 -0
- py_app_runner/bridge/encoders/json_encoder.py +26 -0
- py_app_runner/bridge/encoders/msgpack_encoder.py +58 -0
- py_app_runner/bridge/web_app.py +31 -0
- py_app_runner/bridge/websocket.py +313 -0
- py_app_runner/colors.py +73 -0
- py_app_runner/config.py +132 -0
- py_app_runner/crypto/__init__.py +14 -0
- py_app_runner/crypto/_service.py +75 -0
- py_app_runner/crypto/_service_args.py +54 -0
- py_app_runner/crypto/commands.py +164 -0
- py_app_runner/crypto/envelope.py +144 -0
- py_app_runner/crypto/errors.py +8 -0
- py_app_runner/crypto/fields.py +300 -0
- py_app_runner/crypto/passwords.py +66 -0
- py_app_runner/db_pools.py +20 -0
- py_app_runner/http_exception.py +31 -0
- py_app_runner/logger_handlers.py +167 -0
- py_app_runner/migrations/__init__.py +5 -0
- py_app_runner/migrations/_service.py +296 -0
- py_app_runner/migrations/_service_args.py +91 -0
- py_app_runner/migrations/commands.py +386 -0
- py_app_runner/migrations/discovery.py +108 -0
- py_app_runner/migrations/states.py +63 -0
- py_app_runner/migrations/tracker.py +141 -0
- py_app_runner/py.typed +0 -0
- py_app_runner/pybridge.py +64 -0
- py_app_runner/queue/__init__.py +25 -0
- py_app_runner/queue/_service.py +231 -0
- py_app_runner/queue/_service_args.py +67 -0
- py_app_runner/queue/commands.py +180 -0
- py_app_runner/queue/driver_pg.py +464 -0
- py_app_runner/queue/driver_redis.py +613 -0
- py_app_runner/queue/handler.py +90 -0
- py_app_runner/queue/interface.py +63 -0
- py_app_runner/queue/job.py +46 -0
- py_app_runner/queue/worker.py +221 -0
- py_app_runner/registry.py +54 -0
- py_app_runner/request_handler/__init__.py +0 -0
- py_app_runner/request_handler/auth_service.py +123 -0
- py_app_runner/request_handler/decorators.py +304 -0
- py_app_runner/request_handler/handlers.py +604 -0
- py_app_runner/request_handler/pagination.py +24 -0
- py_app_runner/return_model.py +78 -0
- py_app_runner/runner.py +182 -0
- py_app_runner/throttle/__init__.py +5 -0
- py_app_runner/throttle/throttle.py +217 -0
- py_app_runner/tick_service.py +308 -0
- py_app_runner/timer.py +289 -0
- py_app_runner/utils.py +346 -0
- py_app_runner/wbcm/__init__.py +0 -0
- py_app_runner/wbcm/device_connections.py +89 -0
- py_app_runner/wbcm/factory.py +113 -0
- py_app_runner/wbcm/wb_connection_manager.py +333 -0
- py_app_runner/wbcm/ws_interface.py +56 -0
- py_app_runner-0.5.49.dev0.dist-info/METADATA +134 -0
- py_app_runner-0.5.49.dev0.dist-info/RECORD +75 -0
- py_app_runner-0.5.49.dev0.dist-info/WHEEL +5 -0
- py_app_runner-0.5.49.dev0.dist-info/licenses/LICENSE +21 -0
- py_app_runner-0.5.49.dev0.dist-info/top_level.txt +1 -0
py_app_runner/utils.py
ADDED
|
@@ -0,0 +1,346 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import datetime as dt
|
|
3
|
+
import hashlib
|
|
4
|
+
import json
|
|
5
|
+
import math
|
|
6
|
+
import os
|
|
7
|
+
import re
|
|
8
|
+
import secrets
|
|
9
|
+
import string
|
|
10
|
+
from collections.abc import Callable
|
|
11
|
+
from decimal import ROUND_HALF_UP, Decimal
|
|
12
|
+
from enum import Enum
|
|
13
|
+
from functools import partial
|
|
14
|
+
from typing import Any, ParamSpec, TypeGuard, TypeVar
|
|
15
|
+
|
|
16
|
+
from msgspec import json as mjson
|
|
17
|
+
|
|
18
|
+
P = ParamSpec("P")
|
|
19
|
+
R = TypeVar("R")
|
|
20
|
+
|
|
21
|
+
BASE64_RX = re.compile(r"^[A-Za-z0-9+/=\s]+$")
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def generate_random_string(length: int = 32) -> str:
|
|
25
|
+
"""Generate a random string of fixed length"""
|
|
26
|
+
|
|
27
|
+
return "".join(secrets.choice(string.ascii_letters + string.digits) for _ in range(length))
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
############
|
|
31
|
+
### Json ###
|
|
32
|
+
############
|
|
33
|
+
def json_encoder(obj: Any) -> str | float | None:
|
|
34
|
+
if isinstance(obj, Decimal):
|
|
35
|
+
return float(obj)
|
|
36
|
+
|
|
37
|
+
if isinstance(obj, dt.datetime):
|
|
38
|
+
# The wire format carries no offset, so an aware value must be normalized to
|
|
39
|
+
# UTC first - formatting it as-is would emit a local wall clock time that the
|
|
40
|
+
# client has no way to interpret.
|
|
41
|
+
if obj.tzinfo is not None:
|
|
42
|
+
obj = obj.astimezone(dt.UTC).replace(tzinfo=None)
|
|
43
|
+
return obj.strftime("%Y-%m-%dT%H:%M:%S")
|
|
44
|
+
|
|
45
|
+
if isinstance(obj, dt.date):
|
|
46
|
+
return obj.strftime("%Y-%m-%d")
|
|
47
|
+
|
|
48
|
+
if isinstance(obj, Enum):
|
|
49
|
+
return obj.value
|
|
50
|
+
|
|
51
|
+
if isinstance(obj, float):
|
|
52
|
+
if not math.isfinite(obj):
|
|
53
|
+
return None # "NaN" or "Infinity"
|
|
54
|
+
return obj
|
|
55
|
+
|
|
56
|
+
if isinstance(obj, (int, str)):
|
|
57
|
+
return obj
|
|
58
|
+
|
|
59
|
+
return str(obj)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
class CustomJSONEncoder(json.JSONEncoder):
|
|
63
|
+
"""Stdlib JSONEncoder subclass that delegates to json_encoder for non-standard types."""
|
|
64
|
+
|
|
65
|
+
def default(self, o: Any) -> Any:
|
|
66
|
+
result = json_encoder(o)
|
|
67
|
+
if result is not None:
|
|
68
|
+
return result
|
|
69
|
+
return super().default(o)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
ms_encoder = mjson.Encoder(enc_hook=json_encoder)
|
|
73
|
+
ms_decoder = mjson.Decoder()
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def json_encode_bytes(value: Any, pretty: bool = False) -> bytes:
|
|
77
|
+
json_bytes = ms_encoder.encode(value)
|
|
78
|
+
|
|
79
|
+
if not pretty:
|
|
80
|
+
return json_bytes
|
|
81
|
+
|
|
82
|
+
return mjson.format(json_bytes, indent=2)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def json_encode(value: Any, pretty: bool = False) -> str:
|
|
86
|
+
json_bytes = ms_encoder.encode(value)
|
|
87
|
+
|
|
88
|
+
if not pretty:
|
|
89
|
+
return json_bytes.decode("utf-8")
|
|
90
|
+
|
|
91
|
+
pretty_bytes = mjson.format(json_bytes, indent=2)
|
|
92
|
+
pretty_text = pretty_bytes.decode("utf-8")
|
|
93
|
+
return pretty_text
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def json_decode(value: bytes | str) -> Any:
|
|
97
|
+
return ms_decoder.decode(value)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
###########
|
|
101
|
+
### CPU ###
|
|
102
|
+
###########
|
|
103
|
+
def workers_auto() -> int:
|
|
104
|
+
# honors CPU affinity in containers
|
|
105
|
+
if hasattr(os, "sched_getaffinity"):
|
|
106
|
+
return max(1, len(os.sched_getaffinity(0)))
|
|
107
|
+
return max(1, os.cpu_count() or 1)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
#############
|
|
111
|
+
### Other ###
|
|
112
|
+
#############
|
|
113
|
+
# Ignore exception decorator
|
|
114
|
+
def ignore_exception(IgnoreException: type[Exception]) -> Any:
|
|
115
|
+
"""Decorator for ignoring exception from a function
|
|
116
|
+
e.g. @ignore_exception(DivideByZero)
|
|
117
|
+
e.g.2. ignore_exception(DivideByZero)(Divide)(2/0)
|
|
118
|
+
"""
|
|
119
|
+
|
|
120
|
+
def dec(function: Callable[..., Any]) -> Callable[..., Any]:
|
|
121
|
+
def _dec(*args: tuple[Any] | None, **kwargs: Any | None) -> Any:
|
|
122
|
+
if len(args) > 1:
|
|
123
|
+
newArgs, defaultValue = args[:-1], args[-1]
|
|
124
|
+
else:
|
|
125
|
+
newArgs, defaultValue = args, None
|
|
126
|
+
|
|
127
|
+
try:
|
|
128
|
+
return function(*newArgs, **kwargs)
|
|
129
|
+
except IgnoreException:
|
|
130
|
+
return defaultValue
|
|
131
|
+
|
|
132
|
+
return _dec
|
|
133
|
+
|
|
134
|
+
return dec
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
# Ignore exception if float conversion fails
|
|
138
|
+
sfloat = ignore_exception(ValueError)(float)
|
|
139
|
+
sint = ignore_exception(ValueError)(int)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def is_dict(data: Any) -> TypeGuard[dict[str, Any]]:
|
|
143
|
+
return isinstance(data, dict)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def is_list(data: Any) -> TypeGuard[list[Any]]:
|
|
147
|
+
return isinstance(data, list)
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def pretty_print(data: Any) -> None:
|
|
151
|
+
"""Print a dictionary in a pretty way"""
|
|
152
|
+
|
|
153
|
+
if is_dict(data):
|
|
154
|
+
for key, value in data.items():
|
|
155
|
+
print(f"{key}: {value}")
|
|
156
|
+
|
|
157
|
+
elif is_list(data):
|
|
158
|
+
for item in data:
|
|
159
|
+
pretty_print(item)
|
|
160
|
+
|
|
161
|
+
else:
|
|
162
|
+
print(data)
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
# Fix floats
|
|
166
|
+
def fix_float(string_value: str, default_value: float | None = None) -> float:
|
|
167
|
+
return sfloat(string_value.replace(" ", "").replace(",", ".").strip(), default_value)
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def replace_strings(
|
|
171
|
+
text: str,
|
|
172
|
+
replace_strings: list[tuple[str, str]] | None = None,
|
|
173
|
+
) -> str:
|
|
174
|
+
"""Replace strings in a text"""
|
|
175
|
+
|
|
176
|
+
if replace_strings is not None:
|
|
177
|
+
text = text.strip(' "\r\n') # Remove any extra spaces and quotes
|
|
178
|
+
|
|
179
|
+
for old, new in replace_strings:
|
|
180
|
+
pattern = re.compile(re.escape(old), flags=re.IGNORECASE)
|
|
181
|
+
text = pattern.sub(lambda match, new=new: new, text)
|
|
182
|
+
|
|
183
|
+
text = text.strip(' "\r\n') # Remove any extra spaces and quotes
|
|
184
|
+
|
|
185
|
+
return text
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def create_valid_hostname(input_string: str) -> str:
|
|
189
|
+
if not input_string:
|
|
190
|
+
return ""
|
|
191
|
+
|
|
192
|
+
# Replace spaces with hyphens
|
|
193
|
+
input_string = input_string.replace(" ", "-")
|
|
194
|
+
input_string = input_string.replace("_", "-")
|
|
195
|
+
input_string = input_string.replace(".", "-")
|
|
196
|
+
|
|
197
|
+
# Remove characters that are not letters, digits, or hyphens
|
|
198
|
+
input_string = re.sub(r"[^a-zA-Z0-9-]+", "", input_string)
|
|
199
|
+
|
|
200
|
+
# Ensure that the first and last characters are not hyphens
|
|
201
|
+
input_string = input_string.strip("-")
|
|
202
|
+
|
|
203
|
+
# Return the valid hostname
|
|
204
|
+
return input_string
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def round_decimal(
|
|
208
|
+
value: Decimal,
|
|
209
|
+
digits: int = 0,
|
|
210
|
+
rounding: str = ROUND_HALF_UP,
|
|
211
|
+
) -> Decimal:
|
|
212
|
+
"""
|
|
213
|
+
Rounds a Decimal to the specified number of digits.
|
|
214
|
+
|
|
215
|
+
Args:
|
|
216
|
+
value (Decimal): The Decimal to be rounded.
|
|
217
|
+
digits (int, optional): The number of decimal digits to round to (default is 0).
|
|
218
|
+
rounding (str, optional): The rounding method (default is ROUND_HALF_UP).
|
|
219
|
+
|
|
220
|
+
Returns:
|
|
221
|
+
Decimal: The rounded Decimal.
|
|
222
|
+
"""
|
|
223
|
+
# Calculate the multiplier based on the number of digits
|
|
224
|
+
multiplier = Decimal("10") ** (-digits)
|
|
225
|
+
|
|
226
|
+
# Round the value using quantize with the specified rounding mode
|
|
227
|
+
return value.quantize(multiplier, rounding=rounding)
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
def maybe_truncate(msg: str, maxLen: int = 600) -> tuple[str, str]:
|
|
231
|
+
"""
|
|
232
|
+
Functions that truncates message if it is too long.
|
|
233
|
+
|
|
234
|
+
Always returns tuple, that can be used in logger.debug function call
|
|
235
|
+
as argument. Tuple contains (maybe truncated message, ellipsis)
|
|
236
|
+
"""
|
|
237
|
+
msg_trun = msg[:maxLen]
|
|
238
|
+
ellipsis = ""
|
|
239
|
+
if len(msg) > maxLen:
|
|
240
|
+
ellipsis = "\n...truncated..."
|
|
241
|
+
|
|
242
|
+
return (msg_trun, ellipsis)
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
def safe_decode(raw: Any) -> str:
|
|
246
|
+
if isinstance(raw, (bytes, bytearray)):
|
|
247
|
+
return bytes(raw).decode("utf-8", errors="replace")
|
|
248
|
+
if isinstance(raw, str):
|
|
249
|
+
return raw
|
|
250
|
+
return repr(raw)
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def truncate_middle(s: str, max_len: int) -> str:
|
|
254
|
+
if len(s) <= max_len:
|
|
255
|
+
return s
|
|
256
|
+
head = max_len // 2
|
|
257
|
+
tail = max_len - head - 1
|
|
258
|
+
return f"{s[:head]}…{s[-tail:]}"
|
|
259
|
+
|
|
260
|
+
|
|
261
|
+
def summarize_maybe_base64(s: str, preview: int = 32, hard_cap: int = 256) -> str:
|
|
262
|
+
"""
|
|
263
|
+
Heuristic summary for very long strings. If it looks like base64 (charset + len%4==0),
|
|
264
|
+
show head/tail and approximate decoded size WITHOUT decoding. Otherwise, truncate.
|
|
265
|
+
"""
|
|
266
|
+
compact = "".join(s.split())
|
|
267
|
+
if compact and len(compact) % 4 == 0 and BASE64_RX.match(compact):
|
|
268
|
+
approx = (len(compact) * 3) // 4 - compact.count("=")
|
|
269
|
+
head = compact[:preview]
|
|
270
|
+
tail = compact[-preview:] if len(compact) > 2 * preview else ""
|
|
271
|
+
return f"<base64 {len(compact)} chars (~{approx} bytes) {head}…{tail}>"
|
|
272
|
+
return truncate_middle(s, hard_cap)
|
|
273
|
+
|
|
274
|
+
|
|
275
|
+
async def run_blocking(fn: Callable[P, R], *args: P.args, **kwargs: P.kwargs) -> R:
|
|
276
|
+
"""Run a blocking callable in the default executor without blocking the loop."""
|
|
277
|
+
loop = asyncio.get_running_loop()
|
|
278
|
+
return await loop.run_in_executor(None, partial(fn, *args, **kwargs))
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def redact_headers(headers: dict[str, Any]) -> dict[str, Any]:
|
|
282
|
+
redacted: dict[str, Any] = {}
|
|
283
|
+
for k, v in (headers or {}).items():
|
|
284
|
+
kl = str(k).lower()
|
|
285
|
+
if kl in ("authorization", "proxy-authorization", "cookie", "set-cookie", "x-api-key"):
|
|
286
|
+
redacted[k] = "<redacted>"
|
|
287
|
+
else:
|
|
288
|
+
redacted[k] = v
|
|
289
|
+
return redacted
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
def sha1_prefix(b: bytes, n: int = 12) -> str:
|
|
293
|
+
return hashlib.sha1(b).hexdigest()[:n]
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def sha256_hash(s: str) -> str:
|
|
297
|
+
return hashlib.sha256(s.encode("utf-8")).hexdigest()
|
|
298
|
+
|
|
299
|
+
|
|
300
|
+
def shrink_json(obj: Any, *, max_string: int, max_items: int) -> Any:
|
|
301
|
+
"""
|
|
302
|
+
Traverse JSON-like structures and:
|
|
303
|
+
- truncate any string longer than max_string (base64-aware summary),
|
|
304
|
+
- cap lists to max_items with a "… N more items" marker,
|
|
305
|
+
- leave numbers/bools/None untouched.
|
|
306
|
+
No key-based assumptions.
|
|
307
|
+
"""
|
|
308
|
+
if isinstance(obj, dict):
|
|
309
|
+
return {k: shrink_json(v, max_string=max_string, max_items=max_items) for k, v in obj.items()}
|
|
310
|
+
if isinstance(obj, list):
|
|
311
|
+
trimmed = [shrink_json(x, max_string=max_string, max_items=max_items) for x in obj[:max_items]]
|
|
312
|
+
if len(obj) > max_items:
|
|
313
|
+
trimmed.append(f"<… {len(obj) - max_items} more items>")
|
|
314
|
+
return trimmed
|
|
315
|
+
if isinstance(obj, str):
|
|
316
|
+
return summarize_maybe_base64(obj, hard_cap=max_string)
|
|
317
|
+
return obj
|
|
318
|
+
|
|
319
|
+
|
|
320
|
+
def format_body_for_log(
|
|
321
|
+
raw_body: Any,
|
|
322
|
+
*,
|
|
323
|
+
max_chars: int,
|
|
324
|
+
shrink_strings_to: int,
|
|
325
|
+
cap_list_items: int,
|
|
326
|
+
) -> str:
|
|
327
|
+
text = safe_decode(raw_body)
|
|
328
|
+
try:
|
|
329
|
+
js = json_decode(text)
|
|
330
|
+
js = shrink_json(js, max_string=shrink_strings_to, max_items=cap_list_items)
|
|
331
|
+
return truncate_middle(json_encode(js, pretty=True), max_chars)
|
|
332
|
+
except Exception:
|
|
333
|
+
return truncate_middle(summarize_maybe_base64(text, hard_cap=shrink_strings_to), max_chars)
|
|
334
|
+
|
|
335
|
+
|
|
336
|
+
def format_response_for_log(
|
|
337
|
+
response: Any,
|
|
338
|
+
*,
|
|
339
|
+
max_chars: int,
|
|
340
|
+
shrink_strings_to: int,
|
|
341
|
+
cap_list_items: int,
|
|
342
|
+
) -> str:
|
|
343
|
+
if isinstance(response, (dict, list)):
|
|
344
|
+
js = shrink_json(response, max_string=shrink_strings_to, max_items=cap_list_items)
|
|
345
|
+
return truncate_middle(json_encode(js, pretty=True), max_chars)
|
|
346
|
+
return truncate_middle(safe_decode(response), max_chars)
|
|
File without changes
|
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from collections.abc import ItemsView
|
|
3
|
+
|
|
4
|
+
from py_app_runner.wbcm.ws_interface import WebSocketHandlerInterface
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class DeviceConnections:
|
|
8
|
+
"""Tracks device WebSocket connections, mirroring UserConnections for devices."""
|
|
9
|
+
|
|
10
|
+
# Connection is noted by connection uid
|
|
11
|
+
device_connections: dict[str, WebSocketHandlerInterface]
|
|
12
|
+
|
|
13
|
+
# Map device id to multiple connection uids
|
|
14
|
+
device_id_map: dict[str, list[str]]
|
|
15
|
+
|
|
16
|
+
logger: logging.Logger
|
|
17
|
+
|
|
18
|
+
def __init__(self) -> None:
|
|
19
|
+
self.device_connections = {}
|
|
20
|
+
self.device_id_map = {}
|
|
21
|
+
|
|
22
|
+
logger_name = f"{__name__}.{self.__class__.__name__}"
|
|
23
|
+
self.logger = logging.getLogger(logger_name)
|
|
24
|
+
|
|
25
|
+
def items(self) -> ItemsView[str, WebSocketHandlerInterface]:
|
|
26
|
+
return self.device_connections.items()
|
|
27
|
+
|
|
28
|
+
def add_connection(self, conn: WebSocketHandlerInterface) -> None:
|
|
29
|
+
self.logger.debug("Adding device connection for conn_id %s", conn.uid)
|
|
30
|
+
assert conn.uid, "Connection must have a uid"
|
|
31
|
+
self.device_connections[conn.uid] = conn
|
|
32
|
+
|
|
33
|
+
def remove_connection(self, conn: WebSocketHandlerInterface) -> None:
|
|
34
|
+
self.logger.debug("Removing device connection for conn_id %s", conn.uid)
|
|
35
|
+
assert conn.uid, "Connection must have a uid"
|
|
36
|
+
if conn.uid in self.device_connections:
|
|
37
|
+
del self.device_connections[conn.uid]
|
|
38
|
+
|
|
39
|
+
self.unassign_connection(conn.uid)
|
|
40
|
+
|
|
41
|
+
def unassign_connection(self, conn_id: str, keep_device_id: str | None = None) -> None:
|
|
42
|
+
"""Drop `conn_id` from every device's fan-out list except `keep_device_id`.
|
|
43
|
+
|
|
44
|
+
No early exit: a re-authenticated socket may be present under several ids.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
for device_id, conn_ids in list(self.device_id_map.items()):
|
|
48
|
+
if device_id == keep_device_id:
|
|
49
|
+
continue
|
|
50
|
+
if conn_id in conn_ids:
|
|
51
|
+
conn_ids.remove(conn_id)
|
|
52
|
+
if not conn_ids:
|
|
53
|
+
del self.device_id_map[device_id]
|
|
54
|
+
|
|
55
|
+
def reset(self) -> None:
|
|
56
|
+
self.device_connections = {}
|
|
57
|
+
self.device_id_map = {}
|
|
58
|
+
|
|
59
|
+
def assign_device_id(self, conn_id: str, device_id: str) -> None:
|
|
60
|
+
self.logger.debug("Assigning device_id %s to conn_id %s", device_id, conn_id)
|
|
61
|
+
device_id = str(device_id)
|
|
62
|
+
|
|
63
|
+
# A connection belongs to exactly one device at a time.
|
|
64
|
+
self.unassign_connection(conn_id, keep_device_id=device_id)
|
|
65
|
+
|
|
66
|
+
if device_id not in self.device_id_map:
|
|
67
|
+
self.device_id_map[device_id] = []
|
|
68
|
+
|
|
69
|
+
# Callers may re-assign on every re-auth; keep one entry per connection.
|
|
70
|
+
if conn_id not in self.device_id_map[device_id]:
|
|
71
|
+
self.device_id_map[device_id].append(conn_id)
|
|
72
|
+
|
|
73
|
+
def remove_device_id(self, device_id: str) -> None:
|
|
74
|
+
if device_id in self.device_id_map:
|
|
75
|
+
del self.device_id_map[device_id]
|
|
76
|
+
|
|
77
|
+
def get_connections_by_device_id(
|
|
78
|
+
self,
|
|
79
|
+
device_id: str,
|
|
80
|
+
) -> list[WebSocketHandlerInterface]:
|
|
81
|
+
if device_id not in self.device_id_map:
|
|
82
|
+
return []
|
|
83
|
+
|
|
84
|
+
connections: list[WebSocketHandlerInterface] = []
|
|
85
|
+
for conn_id in self.device_id_map[device_id]:
|
|
86
|
+
if conn_id in self.device_connections:
|
|
87
|
+
connections.append(self.device_connections[conn_id])
|
|
88
|
+
|
|
89
|
+
return connections
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from collections.abc import ItemsView
|
|
3
|
+
from typing import NotRequired, TypedDict
|
|
4
|
+
|
|
5
|
+
from py_app_runner.wbcm.ws_interface import WebSocketHandlerInterface
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class RedisMessage(TypedDict):
|
|
9
|
+
type: str
|
|
10
|
+
service: str
|
|
11
|
+
user_id: str
|
|
12
|
+
source_uid: str
|
|
13
|
+
data: str | bytes | None
|
|
14
|
+
# Only set for `session_revoked` messages targeting a single session.
|
|
15
|
+
target_sid: NotRequired[str]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class UserConnections:
|
|
19
|
+
# Connection is noted by machine id
|
|
20
|
+
user_connections: dict[str, WebSocketHandlerInterface]
|
|
21
|
+
|
|
22
|
+
# Map user id to multiple machine ids
|
|
23
|
+
user_id_map: dict[str, list[str]]
|
|
24
|
+
|
|
25
|
+
# Logger
|
|
26
|
+
logger: logging.Logger
|
|
27
|
+
|
|
28
|
+
def __init__(self) -> None:
|
|
29
|
+
self.user_connections = {}
|
|
30
|
+
self.user_id_map = {}
|
|
31
|
+
|
|
32
|
+
logger_name = f"{__name__}.{self.__class__.__name__}"
|
|
33
|
+
self.logger = logging.getLogger(logger_name)
|
|
34
|
+
|
|
35
|
+
def items(self) -> ItemsView[str, WebSocketHandlerInterface]:
|
|
36
|
+
"""Return a flat list of all connections"""
|
|
37
|
+
|
|
38
|
+
return self.user_connections.items()
|
|
39
|
+
|
|
40
|
+
def add_connection(self, conn: WebSocketHandlerInterface) -> None:
|
|
41
|
+
self.logger.debug(f"Adding connection for conn_id {conn.uid}")
|
|
42
|
+
assert conn.uid, "Connection must have a uid"
|
|
43
|
+
self.user_connections[conn.uid] = conn
|
|
44
|
+
|
|
45
|
+
self.logger.debug(f"Currently we have {len(self.user_connections)} connections")
|
|
46
|
+
|
|
47
|
+
def remove_connection(self, conn: WebSocketHandlerInterface) -> None:
|
|
48
|
+
self.logger.debug(f"Removing connection for conn_id {conn.uid}")
|
|
49
|
+
assert conn.uid, "Connection must have a uid"
|
|
50
|
+
if conn.uid in self.user_connections:
|
|
51
|
+
del self.user_connections[conn.uid]
|
|
52
|
+
|
|
53
|
+
self.unassign_connection(conn.uid)
|
|
54
|
+
|
|
55
|
+
def unassign_connection(self, conn_id: str, keep_user_id: str | None = None) -> None:
|
|
56
|
+
"""Drop `conn_id` from every user's fan-out list except `keep_user_id`.
|
|
57
|
+
|
|
58
|
+
A socket can re-authenticate as a different user mid-connection; without this
|
|
59
|
+
it would stay in the previous user's list and keep receiving that user's
|
|
60
|
+
messages. No early exit: the connection may be present under several ids.
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
for user_id, conn_ids in list(self.user_id_map.items()):
|
|
64
|
+
if user_id == keep_user_id:
|
|
65
|
+
continue
|
|
66
|
+
if conn_id in conn_ids:
|
|
67
|
+
conn_ids.remove(conn_id)
|
|
68
|
+
if not conn_ids:
|
|
69
|
+
del self.user_id_map[user_id]
|
|
70
|
+
|
|
71
|
+
def reset(self) -> None:
|
|
72
|
+
self.logger.debug("Resetting user connections")
|
|
73
|
+
self.user_connections = {}
|
|
74
|
+
self.user_id_map = {}
|
|
75
|
+
|
|
76
|
+
def assign_user_id(self, conn_id: str, user_id: str) -> None:
|
|
77
|
+
self.logger.debug(f"Assigning userId {user_id} to conn_id {conn_id}")
|
|
78
|
+
user_id = str(user_id)
|
|
79
|
+
|
|
80
|
+
# A connection belongs to exactly one user at a time.
|
|
81
|
+
self.unassign_connection(conn_id, keep_user_id=user_id)
|
|
82
|
+
|
|
83
|
+
if user_id not in self.user_id_map:
|
|
84
|
+
self.user_id_map[user_id] = []
|
|
85
|
+
|
|
86
|
+
# Callers may re-assign on every re-auth/subscribe; keep one entry per connection.
|
|
87
|
+
if conn_id not in self.user_id_map[user_id]:
|
|
88
|
+
self.user_id_map[user_id].append(conn_id)
|
|
89
|
+
|
|
90
|
+
self.logger.debug(
|
|
91
|
+
f"Currently we have {len(self.user_id_map)} userIds assigned"
|
|
92
|
+
f" to total of {len(self.user_connections)} connections"
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
def remove_user_id(self, user_id: str) -> None:
|
|
96
|
+
self.logger.debug(f"Removing user_id {user_id}")
|
|
97
|
+
if user_id in self.user_id_map:
|
|
98
|
+
del self.user_id_map[user_id]
|
|
99
|
+
|
|
100
|
+
def get_connections_by_user_id(
|
|
101
|
+
self,
|
|
102
|
+
user_id: str,
|
|
103
|
+
) -> list[WebSocketHandlerInterface]:
|
|
104
|
+
self.logger.debug(f"Getting connections for user_id {user_id}")
|
|
105
|
+
if user_id not in self.user_id_map:
|
|
106
|
+
return []
|
|
107
|
+
|
|
108
|
+
connections: list[WebSocketHandlerInterface] = []
|
|
109
|
+
for conn_id in self.user_id_map[user_id]:
|
|
110
|
+
if conn_id in self.user_connections:
|
|
111
|
+
connections.append(self.user_connections[conn_id])
|
|
112
|
+
|
|
113
|
+
return connections
|