volcano-sdk-python 0.10.0__py3-none-any.whl → 0.10.2__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.
- volcano_sdk/_function_resolution.py +36 -3
- volcano_sdk/_generated/api/projects/get_project_logo.py +21 -10
- volcano_sdk/_generated/api/storage_objects/download_public_file.py +21 -10
- volcano_sdk/_lock_guard.py +10 -3
- volcano_sdk/_lock_renewer.py +9 -1
- volcano_sdk/_log_response.py +84 -0
- volcano_sdk/_realtime_fetch_worker.py +100 -56
- volcano_sdk/_session.py +39 -12
- volcano_sdk/_session_operations.py +9 -1
- volcano_sdk/_transport.py +35 -10
- volcano_sdk/auth.py +286 -77
- volcano_sdk/client.py +16 -6
- volcano_sdk/connection_string.py +15 -1
- volcano_sdk/database.py +217 -40
- volcano_sdk/durable.py +57 -4
- volcano_sdk/durable_authoring.py +159 -48
- volcano_sdk/functions.py +23 -5
- volcano_sdk/locks.py +40 -8
- volcano_sdk/logs.py +27 -36
- volcano_sdk/realtime.py +349 -173
- volcano_sdk/storage.py +87 -15
- {volcano_sdk_python-0.10.0.dist-info → volcano_sdk_python-0.10.2.dist-info}/METADATA +16 -9
- {volcano_sdk_python-0.10.0.dist-info → volcano_sdk_python-0.10.2.dist-info}/RECORD +25 -24
- {volcano_sdk_python-0.10.0.dist-info → volcano_sdk_python-0.10.2.dist-info}/WHEEL +1 -1
- {volcano_sdk_python-0.10.0.dist-info → volcano_sdk_python-0.10.2.dist-info}/licenses/LICENSE +0 -0
|
@@ -28,7 +28,14 @@ _LOCK_STRIPES = 64
|
|
|
28
28
|
|
|
29
29
|
|
|
30
30
|
def _now() -> float:
|
|
31
|
-
"""Read the clock
|
|
31
|
+
"""Read the clock used for cache lifetimes.
|
|
32
|
+
|
|
33
|
+
Returns
|
|
34
|
+
-------
|
|
35
|
+
float
|
|
36
|
+
Monotonic seconds, unaffected by wall-clock adjustments.
|
|
37
|
+
|
|
38
|
+
"""
|
|
32
39
|
return time.monotonic()
|
|
33
40
|
|
|
34
41
|
|
|
@@ -62,7 +69,14 @@ _stripes = [threading.Lock() for _ in range(_LOCK_STRIPES)]
|
|
|
62
69
|
|
|
63
70
|
|
|
64
71
|
def resolve_lock(api_url: str, authorization: str, name: str) -> threading.Lock:
|
|
65
|
-
"""
|
|
72
|
+
"""Select the lock for resolving one name.
|
|
73
|
+
|
|
74
|
+
Returns
|
|
75
|
+
-------
|
|
76
|
+
threading.Lock
|
|
77
|
+
A shared stripe keyed by API URL, credential, and function name.
|
|
78
|
+
|
|
79
|
+
"""
|
|
66
80
|
return _stripes[hash((api_url, authorization, name)) % _LOCK_STRIPES]
|
|
67
81
|
|
|
68
82
|
|
|
@@ -75,6 +89,12 @@ def valid_invoke_url(value: object, api_url: str) -> str | None:
|
|
|
75
89
|
|
|
76
90
|
Anything unusable yields None so the caller falls back to the API path. A
|
|
77
91
|
malformed server response must not raise out of invoke().
|
|
92
|
+
|
|
93
|
+
Returns
|
|
94
|
+
-------
|
|
95
|
+
str or None
|
|
96
|
+
The accepted URL unchanged, or None to use the API endpoint.
|
|
97
|
+
|
|
78
98
|
"""
|
|
79
99
|
if not isinstance(value, str):
|
|
80
100
|
return None
|
|
@@ -91,6 +111,12 @@ def _absolute_url_scheme(value: str) -> str:
|
|
|
91
111
|
|
|
92
112
|
Unparseable input yields "" rather than raising, so a malformed URL reads
|
|
93
113
|
as unusable to every caller.
|
|
114
|
+
|
|
115
|
+
Returns
|
|
116
|
+
-------
|
|
117
|
+
str
|
|
118
|
+
The lowercase scheme, or an empty string for an invalid authority or URL.
|
|
119
|
+
|
|
94
120
|
"""
|
|
95
121
|
if not value or any(character.isspace() for character in value):
|
|
96
122
|
return ""
|
|
@@ -105,7 +131,14 @@ def _absolute_url_scheme(value: str) -> str:
|
|
|
105
131
|
|
|
106
132
|
|
|
107
133
|
def lookup(api_url: str, authorization: str, name: str) -> CachedOutcome | None:
|
|
108
|
-
"""
|
|
134
|
+
"""Look up a function within its API and credential scope.
|
|
135
|
+
|
|
136
|
+
Returns
|
|
137
|
+
-------
|
|
138
|
+
CachedOutcome or None
|
|
139
|
+
An unexpired resolution or remembered miss; None requires a new resolve.
|
|
140
|
+
|
|
141
|
+
"""
|
|
109
142
|
key = (api_url, authorization, name)
|
|
110
143
|
now = _now()
|
|
111
144
|
with _lock:
|
|
@@ -9,6 +9,8 @@ from ...types import Response, UNSET
|
|
|
9
9
|
from ... import errors
|
|
10
10
|
|
|
11
11
|
from ...models.error import Error
|
|
12
|
+
from ...types import File, FileTypes
|
|
13
|
+
from io import BytesIO
|
|
12
14
|
from typing import cast
|
|
13
15
|
from uuid import UUID
|
|
14
16
|
|
|
@@ -34,7 +36,16 @@ def _get_kwargs(
|
|
|
34
36
|
|
|
35
37
|
|
|
36
38
|
|
|
37
|
-
def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Error | None:
|
|
39
|
+
def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Error | File | None:
|
|
40
|
+
if response.status_code == 200:
|
|
41
|
+
response_200 = File(
|
|
42
|
+
payload = BytesIO(response.content)
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
return response_200
|
|
48
|
+
|
|
38
49
|
if response.status_code == 404:
|
|
39
50
|
response_404 = Error.from_dict(response.json())
|
|
40
51
|
|
|
@@ -48,7 +59,7 @@ def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Res
|
|
|
48
59
|
return None
|
|
49
60
|
|
|
50
61
|
|
|
51
|
-
def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Error]:
|
|
62
|
+
def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Error | File]:
|
|
52
63
|
return Response(
|
|
53
64
|
status_code=HTTPStatus(response.status_code),
|
|
54
65
|
content=response.content,
|
|
@@ -62,7 +73,7 @@ def sync_detailed(
|
|
|
62
73
|
*,
|
|
63
74
|
client: AuthenticatedClient | Client,
|
|
64
75
|
|
|
65
|
-
) -> Response[Error]:
|
|
76
|
+
) -> Response[Error | File]:
|
|
66
77
|
""" Get the project logo image
|
|
67
78
|
|
|
68
79
|
Returns the raw logo image stored in the project's storage folder. This
|
|
@@ -79,7 +90,7 @@ def sync_detailed(
|
|
|
79
90
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
80
91
|
|
|
81
92
|
Returns:
|
|
82
|
-
Response[Error]
|
|
93
|
+
Response[Error | File]
|
|
83
94
|
"""
|
|
84
95
|
|
|
85
96
|
|
|
@@ -99,7 +110,7 @@ def sync(
|
|
|
99
110
|
*,
|
|
100
111
|
client: AuthenticatedClient | Client,
|
|
101
112
|
|
|
102
|
-
) -> Error | None:
|
|
113
|
+
) -> Error | File | None:
|
|
103
114
|
""" Get the project logo image
|
|
104
115
|
|
|
105
116
|
Returns the raw logo image stored in the project's storage folder. This
|
|
@@ -116,7 +127,7 @@ def sync(
|
|
|
116
127
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
117
128
|
|
|
118
129
|
Returns:
|
|
119
|
-
Error
|
|
130
|
+
Error | File
|
|
120
131
|
"""
|
|
121
132
|
|
|
122
133
|
|
|
@@ -131,7 +142,7 @@ async def asyncio_detailed(
|
|
|
131
142
|
*,
|
|
132
143
|
client: AuthenticatedClient | Client,
|
|
133
144
|
|
|
134
|
-
) -> Response[Error]:
|
|
145
|
+
) -> Response[Error | File]:
|
|
135
146
|
""" Get the project logo image
|
|
136
147
|
|
|
137
148
|
Returns the raw logo image stored in the project's storage folder. This
|
|
@@ -148,7 +159,7 @@ async def asyncio_detailed(
|
|
|
148
159
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
149
160
|
|
|
150
161
|
Returns:
|
|
151
|
-
Response[Error]
|
|
162
|
+
Response[Error | File]
|
|
152
163
|
"""
|
|
153
164
|
|
|
154
165
|
|
|
@@ -168,7 +179,7 @@ async def asyncio(
|
|
|
168
179
|
*,
|
|
169
180
|
client: AuthenticatedClient | Client,
|
|
170
181
|
|
|
171
|
-
) -> Error | None:
|
|
182
|
+
) -> Error | File | None:
|
|
172
183
|
""" Get the project logo image
|
|
173
184
|
|
|
174
185
|
Returns the raw logo image stored in the project's storage folder. This
|
|
@@ -185,7 +196,7 @@ async def asyncio(
|
|
|
185
196
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
186
197
|
|
|
187
198
|
Returns:
|
|
188
|
-
Error
|
|
199
|
+
Error | File
|
|
189
200
|
"""
|
|
190
201
|
|
|
191
202
|
|
|
@@ -9,6 +9,8 @@ from ...types import Response, UNSET
|
|
|
9
9
|
from ... import errors
|
|
10
10
|
|
|
11
11
|
from ...models.error import Error
|
|
12
|
+
from ...types import File, FileTypes
|
|
13
|
+
from io import BytesIO
|
|
12
14
|
from typing import cast
|
|
13
15
|
from uuid import UUID
|
|
14
16
|
|
|
@@ -36,7 +38,16 @@ def _get_kwargs(
|
|
|
36
38
|
|
|
37
39
|
|
|
38
40
|
|
|
39
|
-
def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Any | Error | None:
|
|
41
|
+
def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Any | Error | File | None:
|
|
42
|
+
if response.status_code == 200:
|
|
43
|
+
response_200 = File(
|
|
44
|
+
payload = BytesIO(response.content)
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
return response_200
|
|
50
|
+
|
|
40
51
|
if response.status_code == 206:
|
|
41
52
|
response_206 = cast(Any, None)
|
|
42
53
|
return response_206
|
|
@@ -58,7 +69,7 @@ def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Res
|
|
|
58
69
|
return None
|
|
59
70
|
|
|
60
71
|
|
|
61
|
-
def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Any | Error]:
|
|
72
|
+
def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Any | Error | File]:
|
|
62
73
|
return Response(
|
|
63
74
|
status_code=HTTPStatus(response.status_code),
|
|
64
75
|
content=response.content,
|
|
@@ -74,7 +85,7 @@ def sync_detailed(
|
|
|
74
85
|
*,
|
|
75
86
|
client: AuthenticatedClient | Client,
|
|
76
87
|
|
|
77
|
-
) -> Response[Any | Error]:
|
|
88
|
+
) -> Response[Any | Error | File]:
|
|
78
89
|
""" Download a public file (no authentication required)
|
|
79
90
|
|
|
80
91
|
Download a file that has been marked as public. This endpoint requires NO authentication.
|
|
@@ -111,7 +122,7 @@ def sync_detailed(
|
|
|
111
122
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
112
123
|
|
|
113
124
|
Returns:
|
|
114
|
-
Response[Any | Error]
|
|
125
|
+
Response[Any | Error | File]
|
|
115
126
|
"""
|
|
116
127
|
|
|
117
128
|
|
|
@@ -135,7 +146,7 @@ def sync(
|
|
|
135
146
|
*,
|
|
136
147
|
client: AuthenticatedClient | Client,
|
|
137
148
|
|
|
138
|
-
) -> Any | Error | None:
|
|
149
|
+
) -> Any | Error | File | None:
|
|
139
150
|
""" Download a public file (no authentication required)
|
|
140
151
|
|
|
141
152
|
Download a file that has been marked as public. This endpoint requires NO authentication.
|
|
@@ -172,7 +183,7 @@ def sync(
|
|
|
172
183
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
173
184
|
|
|
174
185
|
Returns:
|
|
175
|
-
Any | Error
|
|
186
|
+
Any | Error | File
|
|
176
187
|
"""
|
|
177
188
|
|
|
178
189
|
|
|
@@ -191,7 +202,7 @@ async def asyncio_detailed(
|
|
|
191
202
|
*,
|
|
192
203
|
client: AuthenticatedClient | Client,
|
|
193
204
|
|
|
194
|
-
) -> Response[Any | Error]:
|
|
205
|
+
) -> Response[Any | Error | File]:
|
|
195
206
|
""" Download a public file (no authentication required)
|
|
196
207
|
|
|
197
208
|
Download a file that has been marked as public. This endpoint requires NO authentication.
|
|
@@ -228,7 +239,7 @@ async def asyncio_detailed(
|
|
|
228
239
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
229
240
|
|
|
230
241
|
Returns:
|
|
231
|
-
Response[Any | Error]
|
|
242
|
+
Response[Any | Error | File]
|
|
232
243
|
"""
|
|
233
244
|
|
|
234
245
|
|
|
@@ -252,7 +263,7 @@ async def asyncio(
|
|
|
252
263
|
*,
|
|
253
264
|
client: AuthenticatedClient | Client,
|
|
254
265
|
|
|
255
|
-
) -> Any | Error | None:
|
|
266
|
+
) -> Any | Error | File | None:
|
|
256
267
|
""" Download a public file (no authentication required)
|
|
257
268
|
|
|
258
269
|
Download a file that has been marked as public. This endpoint requires NO authentication.
|
|
@@ -289,7 +300,7 @@ async def asyncio(
|
|
|
289
300
|
httpx.TimeoutException: If the request takes longer than Client.timeout.
|
|
290
301
|
|
|
291
302
|
Returns:
|
|
292
|
-
Any | Error
|
|
303
|
+
Any | Error | File
|
|
293
304
|
"""
|
|
294
305
|
|
|
295
306
|
|
volcano_sdk/_lock_guard.py
CHANGED
|
@@ -69,19 +69,26 @@ class LockGuard:
|
|
|
69
69
|
|
|
70
70
|
@property
|
|
71
71
|
def lease(self) -> LockLease:
|
|
72
|
-
"""
|
|
72
|
+
"""The latest successfully renewed immutable lease."""
|
|
73
73
|
with self._state_lock:
|
|
74
74
|
return self._lease
|
|
75
75
|
|
|
76
76
|
@property
|
|
77
77
|
def lost(self) -> bool:
|
|
78
|
-
"""
|
|
78
|
+
"""Whether renewal failed, the lease expired, or the guard closed."""
|
|
79
79
|
with self._state_lock:
|
|
80
80
|
self._expire_if_needed_locked(_lease_now())
|
|
81
81
|
return self._lost.is_set()
|
|
82
82
|
|
|
83
83
|
def wait_lost(self, timeout: float | None = None) -> bool:
|
|
84
|
-
"""Wait until ownership is lost or the timeout expires.
|
|
84
|
+
"""Wait until ownership is lost or the timeout expires.
|
|
85
|
+
|
|
86
|
+
Returns
|
|
87
|
+
-------
|
|
88
|
+
bool
|
|
89
|
+
True when the guard reports lost ownership, False on timeout.
|
|
90
|
+
|
|
91
|
+
"""
|
|
85
92
|
timeout_deadline = None if timeout is None else _lease_now() + timeout
|
|
86
93
|
while True:
|
|
87
94
|
with self._state_lock:
|
volcano_sdk/_lock_renewer.py
CHANGED
|
@@ -17,7 +17,15 @@ def _renewal_jitter() -> float:
|
|
|
17
17
|
|
|
18
18
|
|
|
19
19
|
def renewal_delay(ttl: int, *, remaining: float) -> float:
|
|
20
|
-
"""Choose a jittered renewal time inside the current lease window.
|
|
20
|
+
"""Choose a jittered renewal time inside the current lease window.
|
|
21
|
+
|
|
22
|
+
Returns
|
|
23
|
+
-------
|
|
24
|
+
float
|
|
25
|
+
Seconds before renewal, capped by the remaining safe request window.
|
|
26
|
+
Zero requests immediate renewal.
|
|
27
|
+
|
|
28
|
+
"""
|
|
21
29
|
latest = max(
|
|
22
30
|
0.0,
|
|
23
31
|
remaining - RENEWAL_SAFETY_MARGIN_SECONDS - RENEWAL_REQUEST_BUDGET_SECONDS,
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Validate log response containers before generated model normalization."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from typing import TYPE_CHECKING, cast
|
|
7
|
+
|
|
8
|
+
if TYPE_CHECKING:
|
|
9
|
+
from .models import JSONValue
|
|
10
|
+
|
|
11
|
+
INVALID_LOG_RESPONSE = "Expected a complete log response"
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def response_values(payload: object) -> Mapping[str, object]:
|
|
15
|
+
"""Require a JSON object for the response envelope.
|
|
16
|
+
|
|
17
|
+
Raises:
|
|
18
|
+
TypeError: The response envelope is not an object.
|
|
19
|
+
|
|
20
|
+
Returns:
|
|
21
|
+
The response fields without normalizing their values.
|
|
22
|
+
|
|
23
|
+
"""
|
|
24
|
+
if not isinstance(payload, Mapping):
|
|
25
|
+
raise TypeError(INVALID_LOG_RESPONSE)
|
|
26
|
+
return cast("Mapping[str, object]", payload)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def response_data(values: Mapping[str, object]) -> tuple[Mapping[str, JSONValue], ...]:
|
|
30
|
+
"""Require a list of objects without coercing empty mappings into lists.
|
|
31
|
+
|
|
32
|
+
Raises:
|
|
33
|
+
TypeError: The rows are not a list of objects.
|
|
34
|
+
|
|
35
|
+
Returns:
|
|
36
|
+
The validated rows in their original order.
|
|
37
|
+
|
|
38
|
+
"""
|
|
39
|
+
raw_data = values.get("data")
|
|
40
|
+
if not isinstance(raw_data, list):
|
|
41
|
+
raise TypeError(INVALID_LOG_RESPONSE)
|
|
42
|
+
data = cast("list[object]", raw_data)
|
|
43
|
+
if any(not isinstance(item, Mapping) for item in data):
|
|
44
|
+
raise TypeError(INVALID_LOG_RESPONSE)
|
|
45
|
+
return tuple(cast("Mapping[str, JSONValue]", item) for item in data)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def search_metadata(values: Mapping[str, object]) -> tuple[int, bool, str | None]:
|
|
49
|
+
"""Require search pagination fields without coercing their values.
|
|
50
|
+
|
|
51
|
+
Raises:
|
|
52
|
+
TypeError: Required metadata is missing or has an invalid type.
|
|
53
|
+
|
|
54
|
+
Returns:
|
|
55
|
+
The page limit, continuation flag, and optional cursor.
|
|
56
|
+
|
|
57
|
+
"""
|
|
58
|
+
limit = values.get("limit")
|
|
59
|
+
has_more = values.get("has_more")
|
|
60
|
+
next_cursor = values.get("next_cursor")
|
|
61
|
+
if (
|
|
62
|
+
not isinstance(limit, int)
|
|
63
|
+
or isinstance(limit, bool)
|
|
64
|
+
or not isinstance(has_more, bool)
|
|
65
|
+
or (next_cursor is not None and not isinstance(next_cursor, str))
|
|
66
|
+
):
|
|
67
|
+
raise TypeError(INVALID_LOG_RESPONSE)
|
|
68
|
+
return limit, has_more, next_cursor
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def activity_total(values: Mapping[str, object]) -> int:
|
|
72
|
+
"""Require the activity total without treating booleans as integers.
|
|
73
|
+
|
|
74
|
+
Raises:
|
|
75
|
+
TypeError: The total is missing or is not an integer.
|
|
76
|
+
|
|
77
|
+
Returns:
|
|
78
|
+
The total reported by the server.
|
|
79
|
+
|
|
80
|
+
"""
|
|
81
|
+
total = values.get("total")
|
|
82
|
+
if not isinstance(total, int) or isinstance(total, bool):
|
|
83
|
+
raise TypeError(INVALID_LOG_RESPONSE)
|
|
84
|
+
return total
|
|
@@ -3,16 +3,17 @@
|
|
|
3
3
|
from __future__ import annotations
|
|
4
4
|
|
|
5
5
|
import asyncio
|
|
6
|
+
from collections.abc import Mapping
|
|
6
7
|
from dataclasses import dataclass
|
|
7
|
-
from typing import TYPE_CHECKING,
|
|
8
|
+
from typing import TYPE_CHECKING, Generic, TypeVar
|
|
9
|
+
|
|
10
|
+
from .models import JSONValue
|
|
8
11
|
|
|
9
12
|
if TYPE_CHECKING:
|
|
10
13
|
from collections.abc import Awaitable, Callable
|
|
11
14
|
|
|
12
|
-
from .realtime import _PostgresFetchRequest
|
|
13
|
-
|
|
14
15
|
FallbackT = TypeVar("FallbackT")
|
|
15
|
-
PostgresRecord =
|
|
16
|
+
PostgresRecord = Mapping[str, JSONValue]
|
|
16
17
|
_WORKER_CLOSED = "Postgres fetch worker is closed"
|
|
17
18
|
_INVALID_QUEUE_LIMIT = "queue_limit must be positive"
|
|
18
19
|
_INVALID_BATCH_WINDOW = "batch_window_seconds cannot be negative"
|
|
@@ -20,11 +21,21 @@ _INVALID_BATCH_SIZE = "max_batch_size must be between 1 and queue_limit"
|
|
|
20
21
|
_INVALID_RESULT_COUNT = "Postgres fetch returned an unexpected result count"
|
|
21
22
|
|
|
22
23
|
|
|
24
|
+
@dataclass(frozen=True, slots=True)
|
|
25
|
+
class PostgresFetchRequest:
|
|
26
|
+
"""Capture the database, credential, and row identity for a queued fetch."""
|
|
27
|
+
|
|
28
|
+
database_name: str
|
|
29
|
+
access_token: str
|
|
30
|
+
table: str
|
|
31
|
+
row_id: JSONValue
|
|
32
|
+
|
|
33
|
+
|
|
23
34
|
@dataclass(frozen=True, slots=True)
|
|
24
35
|
class PostgresFetchJob(Generic[FallbackT]):
|
|
25
36
|
"""Pair a captured row request with its lightweight fallback."""
|
|
26
37
|
|
|
27
|
-
request:
|
|
38
|
+
request: PostgresFetchRequest | None
|
|
28
39
|
fallback: FallbackT
|
|
29
40
|
|
|
30
41
|
|
|
@@ -45,13 +56,32 @@ class _StopWorker:
|
|
|
45
56
|
_STOP_WORKER = _StopWorker()
|
|
46
57
|
|
|
47
58
|
|
|
59
|
+
async def _wait_for_close(
|
|
60
|
+
task: asyncio.Task[None],
|
|
61
|
+
stop_task: asyncio.Task[None],
|
|
62
|
+
) -> None:
|
|
63
|
+
completed, _pending = await asyncio.wait(
|
|
64
|
+
(task, stop_task),
|
|
65
|
+
return_when=asyncio.FIRST_COMPLETED,
|
|
66
|
+
)
|
|
67
|
+
if task in completed:
|
|
68
|
+
try:
|
|
69
|
+
task.result()
|
|
70
|
+
except BaseException:
|
|
71
|
+
_ = stop_task.cancel()
|
|
72
|
+
_ = await asyncio.gather(stop_task, return_exceptions=True)
|
|
73
|
+
raise
|
|
74
|
+
await asyncio.shield(stop_task)
|
|
75
|
+
await asyncio.shield(task)
|
|
76
|
+
|
|
77
|
+
|
|
48
78
|
class PostgresFetchWorker(Generic[FallbackT]):
|
|
49
79
|
"""Fetch queued rows serially and deliver outcomes in enqueue order."""
|
|
50
80
|
|
|
51
81
|
def __init__(
|
|
52
82
|
self,
|
|
53
83
|
fetch: Callable[
|
|
54
|
-
[tuple[
|
|
84
|
+
[tuple[PostgresFetchRequest, ...]],
|
|
55
85
|
Awaitable[tuple[PostgresRecord | None, ...]],
|
|
56
86
|
],
|
|
57
87
|
deliver: Callable[[PostgresFetchOutcome[FallbackT]], Awaitable[None]],
|
|
@@ -60,27 +90,47 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
60
90
|
batch_window_seconds: float = 0,
|
|
61
91
|
max_batch_size: int = 1,
|
|
62
92
|
) -> None:
|
|
63
|
-
"""Create a worker with a fixed pending-job limit.
|
|
93
|
+
"""Create a worker with a fixed pending-job limit.
|
|
94
|
+
|
|
95
|
+
Raises
|
|
96
|
+
------
|
|
97
|
+
ValueError
|
|
98
|
+
The queue limit is not positive, the batch window is negative,
|
|
99
|
+
or the batch size is outside the range from one to the queue limit.
|
|
100
|
+
|
|
101
|
+
"""
|
|
64
102
|
if queue_limit <= 0:
|
|
65
103
|
raise ValueError(_INVALID_QUEUE_LIMIT)
|
|
66
104
|
if batch_window_seconds < 0:
|
|
67
105
|
raise ValueError(_INVALID_BATCH_WINDOW)
|
|
68
106
|
if not 1 <= max_batch_size <= queue_limit:
|
|
69
107
|
raise ValueError(_INVALID_BATCH_SIZE)
|
|
70
|
-
self._fetch
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
108
|
+
self._fetch: Callable[
|
|
109
|
+
[tuple[PostgresFetchRequest, ...]],
|
|
110
|
+
Awaitable[tuple[PostgresRecord | None, ...]],
|
|
111
|
+
] = fetch
|
|
112
|
+
self._deliver: Callable[[PostgresFetchOutcome[FallbackT]], Awaitable[None]] = (
|
|
113
|
+
deliver
|
|
114
|
+
)
|
|
115
|
+
self._batch_window_seconds: float = batch_window_seconds
|
|
116
|
+
self._max_batch_size: int = max_batch_size
|
|
74
117
|
self._queue: asyncio.Queue[PostgresFetchJob[FallbackT] | _StopWorker] = (
|
|
75
118
|
asyncio.Queue(maxsize=queue_limit)
|
|
76
119
|
)
|
|
77
|
-
self._state_lock = asyncio.Lock()
|
|
120
|
+
self._state_lock: asyncio.Lock = asyncio.Lock()
|
|
78
121
|
self._task: asyncio.Task[None] | None = None
|
|
79
122
|
self._stop_task: asyncio.Task[None] | None = None
|
|
80
|
-
self._closed = False
|
|
123
|
+
self._closed: bool = False
|
|
81
124
|
|
|
82
125
|
async def enqueue(self, job: PostgresFetchJob[FallbackT]) -> None:
|
|
83
|
-
"""Queue one fetch, applying backpressure when the queue is full.
|
|
126
|
+
"""Queue one fetch, applying backpressure when the queue is full.
|
|
127
|
+
|
|
128
|
+
Raises
|
|
129
|
+
------
|
|
130
|
+
RuntimeError
|
|
131
|
+
The worker has closed or was cancelled before accepting the job.
|
|
132
|
+
|
|
133
|
+
"""
|
|
84
134
|
async with self._state_lock:
|
|
85
135
|
if self._closed:
|
|
86
136
|
raise RuntimeError(_WORKER_CLOSED)
|
|
@@ -102,7 +152,7 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
102
152
|
self._stop_task = asyncio.create_task(self._queue.put(_STOP_WORKER))
|
|
103
153
|
stop_task = self._stop_task
|
|
104
154
|
if task is not None and stop_task is not None:
|
|
105
|
-
await
|
|
155
|
+
await _wait_for_close(task, stop_task)
|
|
106
156
|
|
|
107
157
|
async def abort(self) -> None:
|
|
108
158
|
"""Discard obsolete jobs and stop without waiting for row fetches."""
|
|
@@ -110,14 +160,14 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
110
160
|
task = self._task
|
|
111
161
|
stop_task = self._stop_task
|
|
112
162
|
if stop_task is not None and not stop_task.done():
|
|
113
|
-
stop_task.cancel()
|
|
163
|
+
_ = stop_task.cancel()
|
|
114
164
|
if task is not None and not task.done():
|
|
115
|
-
task.cancel()
|
|
165
|
+
_ = task.cancel()
|
|
116
166
|
pending = tuple(
|
|
117
167
|
candidate for candidate in (task, stop_task) if candidate is not None
|
|
118
168
|
)
|
|
119
169
|
if pending:
|
|
120
|
-
await asyncio.gather(*pending, return_exceptions=True)
|
|
170
|
+
_ = await asyncio.gather(*pending, return_exceptions=True)
|
|
121
171
|
self._discard_pending()
|
|
122
172
|
|
|
123
173
|
def _raise_worker_failure(self) -> None:
|
|
@@ -138,8 +188,8 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
138
188
|
)
|
|
139
189
|
finally:
|
|
140
190
|
if not put_task.done():
|
|
141
|
-
put_task.cancel()
|
|
142
|
-
await asyncio.gather(put_task, return_exceptions=True)
|
|
191
|
+
_ = put_task.cancel()
|
|
192
|
+
_ = await asyncio.gather(put_task, return_exceptions=True)
|
|
143
193
|
if task in completed:
|
|
144
194
|
if task.cancelled():
|
|
145
195
|
raise RuntimeError(_WORKER_CLOSED)
|
|
@@ -148,28 +198,9 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
148
198
|
|
|
149
199
|
def _discard_pending(self) -> None:
|
|
150
200
|
while not self._queue.empty():
|
|
151
|
-
self._queue.get_nowait()
|
|
201
|
+
_ = self._queue.get_nowait()
|
|
152
202
|
self._queue.task_done()
|
|
153
203
|
|
|
154
|
-
async def _wait_for_close(
|
|
155
|
-
self,
|
|
156
|
-
task: asyncio.Task[None],
|
|
157
|
-
stop_task: asyncio.Task[None],
|
|
158
|
-
) -> None:
|
|
159
|
-
completed, _pending = await asyncio.wait(
|
|
160
|
-
(task, stop_task),
|
|
161
|
-
return_when=asyncio.FIRST_COMPLETED,
|
|
162
|
-
)
|
|
163
|
-
if task in completed:
|
|
164
|
-
try:
|
|
165
|
-
task.result()
|
|
166
|
-
except BaseException:
|
|
167
|
-
stop_task.cancel()
|
|
168
|
-
await asyncio.gather(stop_task, return_exceptions=True)
|
|
169
|
-
raise
|
|
170
|
-
await asyncio.shield(stop_task)
|
|
171
|
-
await asyncio.shield(task)
|
|
172
|
-
|
|
173
204
|
async def _run(self) -> None:
|
|
174
205
|
pending: PostgresFetchJob[FallbackT] | _StopWorker | None = None
|
|
175
206
|
while True:
|
|
@@ -196,15 +227,8 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
196
227
|
return None
|
|
197
228
|
deadline = asyncio.get_running_loop().time() + self._batch_window_seconds
|
|
198
229
|
while len(batch) < self._max_batch_size:
|
|
199
|
-
|
|
200
|
-
if
|
|
201
|
-
return None
|
|
202
|
-
try:
|
|
203
|
-
async with asyncio.timeout(remaining):
|
|
204
|
-
candidate = await self._queue.get()
|
|
205
|
-
except TimeoutError:
|
|
206
|
-
return None
|
|
207
|
-
if isinstance(candidate, _StopWorker):
|
|
230
|
+
candidate = await self._next_before(deadline)
|
|
231
|
+
if candidate is None or isinstance(candidate, _StopWorker):
|
|
208
232
|
return candidate
|
|
209
233
|
request = candidate.request
|
|
210
234
|
if (
|
|
@@ -216,10 +240,23 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
216
240
|
batch.append(candidate)
|
|
217
241
|
return None
|
|
218
242
|
|
|
243
|
+
async def _next_before(
|
|
244
|
+
self,
|
|
245
|
+
deadline: float,
|
|
246
|
+
) -> PostgresFetchJob[FallbackT] | _StopWorker | None:
|
|
247
|
+
remaining = deadline - asyncio.get_running_loop().time()
|
|
248
|
+
if remaining <= 0:
|
|
249
|
+
return None
|
|
250
|
+
try:
|
|
251
|
+
async with asyncio.timeout(remaining):
|
|
252
|
+
return await self._queue.get()
|
|
253
|
+
except TimeoutError:
|
|
254
|
+
return None
|
|
255
|
+
|
|
219
256
|
@staticmethod
|
|
220
257
|
def _same_fetch(
|
|
221
|
-
first:
|
|
222
|
-
request:
|
|
258
|
+
first: PostgresFetchRequest,
|
|
259
|
+
request: PostgresFetchRequest,
|
|
223
260
|
) -> bool:
|
|
224
261
|
return (
|
|
225
262
|
first.database_name,
|
|
@@ -234,7 +271,7 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
234
271
|
@staticmethod
|
|
235
272
|
def _repeats_row(
|
|
236
273
|
batch: list[PostgresFetchJob[FallbackT]],
|
|
237
|
-
request:
|
|
274
|
+
request: PostgresFetchRequest,
|
|
238
275
|
) -> bool:
|
|
239
276
|
return any(
|
|
240
277
|
job.request is not None and job.request.row_id == request.row_id
|
|
@@ -254,12 +291,19 @@ class PostgresFetchWorker(Generic[FallbackT]):
|
|
|
254
291
|
return_exceptions=True,
|
|
255
292
|
)
|
|
256
293
|
if isinstance(result, BaseException):
|
|
257
|
-
|
|
258
|
-
raise result
|
|
259
|
-
for job in jobs:
|
|
260
|
-
await self._deliver(PostgresFetchOutcome(job=job, error=result))
|
|
294
|
+
await self._deliver_failure(jobs, result)
|
|
261
295
|
return
|
|
262
296
|
if len(result) != len(jobs):
|
|
263
297
|
raise RuntimeError(_INVALID_RESULT_COUNT)
|
|
264
298
|
for job, record in zip(jobs, result, strict=True):
|
|
265
299
|
await self._deliver(PostgresFetchOutcome(job=job, record=record))
|
|
300
|
+
|
|
301
|
+
async def _deliver_failure(
|
|
302
|
+
self,
|
|
303
|
+
jobs: list[PostgresFetchJob[FallbackT]],
|
|
304
|
+
error: BaseException,
|
|
305
|
+
) -> None:
|
|
306
|
+
if not isinstance(error, Exception):
|
|
307
|
+
raise error
|
|
308
|
+
for job in jobs:
|
|
309
|
+
await self._deliver(PostgresFetchOutcome(job=job, error=error))
|