hyperforge-nucliadb 1.0.0.post99__tar.gz → 1.0.0.post105__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.
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/PKG-INFO +1 -1
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/pyproject.toml +1 -1
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/sync/agent.py +189 -23
- hyperforge_nucliadb-1.0.0.post105/src/hyperforge_nucliadb/sync/config_driver.py +40 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/sync/driver.py +33 -15
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb.egg-info/PKG-INFO +1 -1
- hyperforge_nucliadb-1.0.0.post105/tests/test_sync.py +573 -0
- hyperforge_nucliadb-1.0.0.post99/src/hyperforge_nucliadb/sync/config_driver.py +0 -25
- hyperforge_nucliadb-1.0.0.post99/tests/test_sync.py +0 -223
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/README.md +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/setup.cfg +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/__init__.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/advanced_ask_agent.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/advanced_ask_config.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/__init__.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/analysis.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/ask.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/config.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/hydrate.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/kb_analysis.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/knowledge_scan.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/models.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/multi.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/nucliadb.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/prompt_analysis.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/query_analysis.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/rerank.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask/utils.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/ask_utils.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/basic_ask_agent.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/basic_ask_config.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/driver.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/driver_config.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/py.typed +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/sync/__init__.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb/sync/config.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb.egg-info/SOURCES.txt +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb.egg-info/dependency_links.txt +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb.egg-info/requires.txt +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/src/hyperforge_nucliadb.egg-info/top_level.txt +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/tests/test_ask_utils.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/tests/test_driver.py +0 -0
- {hyperforge_nucliadb-1.0.0.post99 → hyperforge_nucliadb-1.0.0.post105}/tests/test_nucliadb.py +0 -0
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "hyperforge_nucliadb"
|
|
7
|
-
version = "1.0.0.
|
|
7
|
+
version = "1.0.0.post105"
|
|
8
8
|
license = "Apache-2.0"
|
|
9
9
|
description = "NucliaDB Hyperforge agent"
|
|
10
10
|
authors = [{ name = "Nuclia", email = "nucliadb@nuclia.com" }]
|
|
@@ -3,6 +3,7 @@ import itertools
|
|
|
3
3
|
from typing import Dict, List, Optional, cast
|
|
4
4
|
from uuid import uuid4
|
|
5
5
|
|
|
6
|
+
from hyperforge import logger
|
|
6
7
|
from hyperforge.configure import agent
|
|
7
8
|
from hyperforge.interaction import (
|
|
8
9
|
Feedback,
|
|
@@ -109,13 +110,159 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
109
110
|
memory: QuestionMemory,
|
|
110
111
|
manager: Manager,
|
|
111
112
|
resources: Dict[str, List[str]],
|
|
113
|
+
) -> Dict[str, Dict[str, List[str]]]:
|
|
114
|
+
connections = self.get_connections(manager)
|
|
115
|
+
filtered_resources: Dict[str, Dict[str, List[str]]] = {}
|
|
116
|
+
for kb_source_id, resource_ids in resources.items():
|
|
117
|
+
if connections[kb_source_id].config.connection_ids:
|
|
118
|
+
source_resources = await self._post_filter_configured_resources(
|
|
119
|
+
memory=memory,
|
|
120
|
+
manager=manager,
|
|
121
|
+
resources={kb_source_id: resource_ids},
|
|
122
|
+
)
|
|
123
|
+
else:
|
|
124
|
+
source_resources = await self._post_filter_hybrid_resources(
|
|
125
|
+
memory=memory,
|
|
126
|
+
manager=manager,
|
|
127
|
+
kb_source_id=kb_source_id,
|
|
128
|
+
resource_ids=resource_ids,
|
|
129
|
+
)
|
|
130
|
+
filtered_resources.update(source_resources)
|
|
131
|
+
return filtered_resources
|
|
132
|
+
|
|
133
|
+
async def _post_filter_hybrid_resources(
|
|
134
|
+
self,
|
|
135
|
+
memory: QuestionMemory,
|
|
136
|
+
manager: Manager,
|
|
137
|
+
kb_source_id: str,
|
|
138
|
+
resource_ids: List[str],
|
|
139
|
+
) -> Dict[str, Dict[str, List[str]]]:
|
|
140
|
+
driver = self.sources[kb_source_id]
|
|
141
|
+
ndb = get_ndb_driver(manager, kb_source_id)
|
|
142
|
+
resource_objs: List[NucliaDBResource] = await asyncio.gather(
|
|
143
|
+
*[
|
|
144
|
+
ndb.driver.get_resource_by_id(
|
|
145
|
+
kbid=ndb.config.kbid,
|
|
146
|
+
rid=resource_id,
|
|
147
|
+
query_params={"show": [ResourceProperties.ORIGIN.value]},
|
|
148
|
+
)
|
|
149
|
+
for resource_id in resource_ids
|
|
150
|
+
]
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
public_resource_ids: List[str] = []
|
|
154
|
+
connected_resources: Dict[str, List[str]] = {}
|
|
155
|
+
for resource_obj in resource_objs:
|
|
156
|
+
origin = resource_obj.origin
|
|
157
|
+
if origin is None:
|
|
158
|
+
public_resource_ids.append(resource_obj.id)
|
|
159
|
+
continue
|
|
160
|
+
|
|
161
|
+
if not origin.source_id or not origin.source_id.startswith("sync_config_"):
|
|
162
|
+
if origin.sync_metadata is None:
|
|
163
|
+
public_resource_ids.append(resource_obj.id)
|
|
164
|
+
else:
|
|
165
|
+
logger.warning(
|
|
166
|
+
"Excluding resource with inconsistent Sync origin metadata",
|
|
167
|
+
extra={"resource_id": resource_obj.id},
|
|
168
|
+
)
|
|
169
|
+
continue
|
|
170
|
+
|
|
171
|
+
source_id = origin.source_id
|
|
172
|
+
sync_config_id = source_id.removeprefix("sync_config_")
|
|
173
|
+
if not sync_config_id or origin.sync_metadata is None:
|
|
174
|
+
logger.warning(
|
|
175
|
+
"Excluding synced resource with incomplete origin metadata",
|
|
176
|
+
extra={"resource_id": resource_obj.id, "source_id": source_id},
|
|
177
|
+
)
|
|
178
|
+
continue
|
|
179
|
+
connected_resources.setdefault(sync_config_id, []).append(resource_obj.id)
|
|
180
|
+
|
|
181
|
+
resolved_resource_ids: List[str] = []
|
|
182
|
+
resolved_connection_ids: set[str] = set()
|
|
183
|
+
for sync_config_id, connection_resource_ids in connected_resources.items():
|
|
184
|
+
try:
|
|
185
|
+
await driver.resolve_sync_config(sync_config_id)
|
|
186
|
+
except Exception as exc:
|
|
187
|
+
logger.warning(
|
|
188
|
+
"Excluding resources for an unavailable sync configuration",
|
|
189
|
+
extra={
|
|
190
|
+
"sync_config_id": sync_config_id,
|
|
191
|
+
"resource_count": len(connection_resource_ids),
|
|
192
|
+
"exception_type": type(exc).__name__,
|
|
193
|
+
},
|
|
194
|
+
)
|
|
195
|
+
continue
|
|
196
|
+
resolved_resource_ids.extend(connection_resource_ids)
|
|
197
|
+
resolved_connection_ids.add(sync_config_id)
|
|
198
|
+
|
|
199
|
+
public_resource_id_set = set(public_resource_ids)
|
|
200
|
+
if not resolved_resource_ids:
|
|
201
|
+
return {
|
|
202
|
+
kb_source_id: {
|
|
203
|
+
"__hybrid__": [
|
|
204
|
+
resource_id
|
|
205
|
+
for resource_id in resource_ids
|
|
206
|
+
if resource_id in public_resource_id_set
|
|
207
|
+
]
|
|
208
|
+
}
|
|
209
|
+
}
|
|
210
|
+
|
|
211
|
+
try:
|
|
212
|
+
connected_filtered_resources = await self._post_filter_configured_resources(
|
|
213
|
+
memory=memory,
|
|
214
|
+
manager=manager,
|
|
215
|
+
resources={kb_source_id: resolved_resource_ids},
|
|
216
|
+
allowed_connection_ids=resolved_connection_ids,
|
|
217
|
+
)
|
|
218
|
+
except Exception as exc:
|
|
219
|
+
logger.warning(
|
|
220
|
+
"Excluding synced resources after authorization failed",
|
|
221
|
+
extra={
|
|
222
|
+
"source_id": kb_source_id,
|
|
223
|
+
"exception_type": type(exc).__name__,
|
|
224
|
+
},
|
|
225
|
+
)
|
|
226
|
+
connected_filtered_resources = {}
|
|
227
|
+
|
|
228
|
+
authorized_resource_ids = {
|
|
229
|
+
resource_id
|
|
230
|
+
for connection_resource_ids in connected_filtered_resources.get(
|
|
231
|
+
kb_source_id, {}
|
|
232
|
+
).values()
|
|
233
|
+
for resource_id in connection_resource_ids
|
|
234
|
+
}
|
|
235
|
+
return {
|
|
236
|
+
kb_source_id: {
|
|
237
|
+
"__hybrid__": [
|
|
238
|
+
resource_id
|
|
239
|
+
for resource_id in resource_ids
|
|
240
|
+
if resource_id in public_resource_id_set
|
|
241
|
+
or resource_id in authorized_resource_ids
|
|
242
|
+
]
|
|
243
|
+
}
|
|
244
|
+
}
|
|
245
|
+
|
|
246
|
+
async def _post_filter_configured_resources(
|
|
247
|
+
self,
|
|
248
|
+
memory: QuestionMemory,
|
|
249
|
+
manager: Manager,
|
|
250
|
+
resources: Dict[str, List[str]],
|
|
251
|
+
allowed_connection_ids: Optional[set[str]] = None,
|
|
112
252
|
) -> Dict[str, Dict[str, List[str]]]:
|
|
113
253
|
"""We have a list of ARAG resources per KB source. Now we need to filter them by connection"""
|
|
114
254
|
filtered_resources: Dict[str, Dict[str, List[str]]] = {}
|
|
115
|
-
connections_by_resource: Dict[str, List[str]] = {}
|
|
116
|
-
sync_metadata_by_resource: Dict[str, SyncMetadata] = {}
|
|
117
255
|
connections = self.get_connections(manager)
|
|
118
256
|
for kb_source_id, resource_ids in resources.items():
|
|
257
|
+
driver = self.sources[kb_source_id]
|
|
258
|
+
configured_connection_ids = (
|
|
259
|
+
allowed_connection_ids
|
|
260
|
+
if allowed_connection_ids is not None
|
|
261
|
+
else set(driver.config.connection_ids)
|
|
262
|
+
)
|
|
263
|
+
connections_by_resource: Dict[str, List[str]] = {}
|
|
264
|
+
sync_metadata_by_resource: Dict[str, SyncMetadata] = {}
|
|
265
|
+
|
|
119
266
|
# Get resource source connection for each resource
|
|
120
267
|
ndb = get_ndb_driver(manager, kb_source_id)
|
|
121
268
|
get_resource_ids_tasks = []
|
|
@@ -131,24 +278,28 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
131
278
|
*get_resource_ids_tasks
|
|
132
279
|
)
|
|
133
280
|
for resource_obj in resource_objs:
|
|
281
|
+
origin = resource_obj.origin
|
|
134
282
|
if (
|
|
135
|
-
|
|
136
|
-
|
|
283
|
+
origin is None
|
|
284
|
+
or not origin.source_id
|
|
285
|
+
or not origin.source_id.startswith("sync_config_")
|
|
286
|
+
or origin.sync_metadata is None
|
|
137
287
|
):
|
|
138
|
-
|
|
139
|
-
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
|
|
146
|
-
|
|
147
|
-
|
|
288
|
+
continue
|
|
289
|
+
|
|
290
|
+
source_id = origin.source_id.removeprefix("sync_config_")
|
|
291
|
+
if source_id not in configured_connection_ids:
|
|
292
|
+
continue
|
|
293
|
+
|
|
294
|
+
connections_by_resource.setdefault(source_id, []).append(
|
|
295
|
+
resource_obj.id
|
|
296
|
+
)
|
|
297
|
+
sync_metadata_by_resource[resource_obj.id] = origin.sync_metadata
|
|
148
298
|
|
|
149
299
|
# Get providers needed for the resources allocated
|
|
150
300
|
needed_providers_ids = list(connections_by_resource.keys())
|
|
151
|
-
|
|
301
|
+
if not needed_providers_ids:
|
|
302
|
+
continue
|
|
152
303
|
|
|
153
304
|
creds_providers = {}
|
|
154
305
|
for sync_config_id in needed_providers_ids:
|
|
@@ -258,13 +409,19 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
258
409
|
raise Exception("Connection not found")
|
|
259
410
|
|
|
260
411
|
credential = connection_credentials[inner_connection_id]
|
|
412
|
+
connection_resource_ids = connections_by_resource[connection_id]
|
|
413
|
+
connection_sync_metadata = {
|
|
414
|
+
resource_id: sync_metadata_by_resource[resource_id]
|
|
415
|
+
for resource_id in connection_resource_ids
|
|
416
|
+
if resource_id in sync_metadata_by_resource
|
|
417
|
+
}
|
|
261
418
|
filtered_resource_ids = await self.validate_resources_by_connection(
|
|
262
419
|
connection=source,
|
|
263
420
|
credentials=credential,
|
|
264
|
-
resource_ids=
|
|
421
|
+
resource_ids=connection_resource_ids,
|
|
265
422
|
connection_id=inner_connection_id,
|
|
266
423
|
sync_config_id=connection_id,
|
|
267
|
-
sync_metadata_by_resource=
|
|
424
|
+
sync_metadata_by_resource=connection_sync_metadata,
|
|
268
425
|
)
|
|
269
426
|
|
|
270
427
|
filtered_resources.setdefault(kb_source_id, {})[connection_id] = (
|
|
@@ -279,13 +436,24 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
279
436
|
self,
|
|
280
437
|
catalog_filter: Optional[CatalogFilterExpression] = None,
|
|
281
438
|
):
|
|
439
|
+
if any(
|
|
440
|
+
not self.sources[source].config.connection_ids
|
|
441
|
+
for source in self.config.sources
|
|
442
|
+
):
|
|
443
|
+
return catalog_filter
|
|
444
|
+
|
|
445
|
+
connection_ids = [
|
|
446
|
+
connection_id
|
|
447
|
+
for source in self.config.sources
|
|
448
|
+
for connection_id in self.sources[source].config.connection_ids
|
|
449
|
+
]
|
|
450
|
+
|
|
282
451
|
if catalog_filter is None:
|
|
283
452
|
catalog_filter = CatalogFilterExpression(
|
|
284
453
|
resource=Or(
|
|
285
454
|
operands=[
|
|
286
455
|
OriginSource(id=f"sync_config_{connection_id}")
|
|
287
|
-
for
|
|
288
|
-
for connection_id in self.sources[source].sync_configs.keys()
|
|
456
|
+
for connection_id in connection_ids
|
|
289
457
|
]
|
|
290
458
|
)
|
|
291
459
|
)
|
|
@@ -296,10 +464,7 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
296
464
|
Or(
|
|
297
465
|
operands=[
|
|
298
466
|
OriginSource(id=f"sync_config_{connection_id}")
|
|
299
|
-
for
|
|
300
|
-
for connection_id in self.sources[
|
|
301
|
-
source
|
|
302
|
-
].sync_configs.keys()
|
|
467
|
+
for connection_id in connection_ids
|
|
303
468
|
]
|
|
304
469
|
),
|
|
305
470
|
]
|
|
@@ -316,6 +481,7 @@ class SyncAskAgent(BasicAskAgent):
|
|
|
316
481
|
catalog_filter: Optional[CatalogFilterExpression] = None,
|
|
317
482
|
**kwargs,
|
|
318
483
|
) -> Dict[str, List[str]]:
|
|
484
|
+
self.get_connections(manager)
|
|
319
485
|
catalog_filter = self.enrich_catalog_filter(catalog_filter)
|
|
320
486
|
resources = await super().search_by_title(
|
|
321
487
|
memory=memory,
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
from typing import Literal
|
|
2
|
+
|
|
3
|
+
from httpx import Timeout
|
|
4
|
+
from hyperforge.driver import DriverConfig
|
|
5
|
+
from pydantic import Field, field_validator
|
|
6
|
+
from pydantic.config import ConfigDict
|
|
7
|
+
|
|
8
|
+
from hyperforge_nucliadb.driver_config import (
|
|
9
|
+
NucliaDBConnection,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
SYNC_HTTP_TIMEOUT = Timeout(connect=5.0, read=30.0, write=10.0, pool=5.0)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class SyncConnection(NucliaDBConnection):
|
|
16
|
+
connection_ids: list[str] = Field(
|
|
17
|
+
default_factory=list,
|
|
18
|
+
description=(
|
|
19
|
+
"Sync configuration IDs to restrict searches to. When omitted or empty, "
|
|
20
|
+
"the entire knowledge box is searched and synced resources are authorized "
|
|
21
|
+
"individually."
|
|
22
|
+
),
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
@field_validator("connection_ids")
|
|
26
|
+
@classmethod
|
|
27
|
+
def validate_connection_ids(cls, connection_ids: list[str]) -> list[str]:
|
|
28
|
+
if any(not connection_id.strip() for connection_id in connection_ids):
|
|
29
|
+
raise ValueError("connection_ids cannot contain blank values")
|
|
30
|
+
return connection_ids
|
|
31
|
+
|
|
32
|
+
@property
|
|
33
|
+
def kb_url(self) -> str:
|
|
34
|
+
return f"{self.url}/v1/kb/{self.kbid}"
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class SyncDriverConfig(DriverConfig[SyncConnection]):
|
|
38
|
+
model_config = ConfigDict(title="Knowledge Box Sync Service connection")
|
|
39
|
+
provider: Literal["sync"]
|
|
40
|
+
config: SyncConnection
|
|
@@ -90,29 +90,47 @@ class SyncDriver(NucliaDBDriver):
|
|
|
90
90
|
async def init(cls, driver: Any) -> "SyncDriver":
|
|
91
91
|
sync_driver = cast(SyncDriverConfig, driver)
|
|
92
92
|
client = await sync_connect(sync_driver.config)
|
|
93
|
-
|
|
94
|
-
sync_configs: Dict[str, list[str]] = {}
|
|
95
|
-
for sync_config_id in sync_driver.config.connection_ids:
|
|
96
|
-
response = await client.get(f"/sync_config/{sync_config_id}")
|
|
97
|
-
response.raise_for_status()
|
|
98
|
-
data = response.json()
|
|
99
|
-
connection_id = data["external_connection"]["id"]
|
|
100
|
-
response = await client.get(f"/external_connection/{connection_id}")
|
|
101
|
-
response.raise_for_status()
|
|
102
|
-
sync_configs.setdefault(sync_config_id, []).append(connection_id)
|
|
103
|
-
data = response.json()
|
|
104
|
-
information[connection_id] = ExternalConnectionOutput(**data)
|
|
105
|
-
return cls(
|
|
93
|
+
instance = cls(
|
|
106
94
|
provider=sync_driver.provider,
|
|
107
95
|
name=sync_driver.name,
|
|
108
96
|
async_driver=client,
|
|
109
|
-
information=
|
|
110
|
-
sync_configs=
|
|
97
|
+
information={},
|
|
98
|
+
sync_configs={},
|
|
111
99
|
config=sync_driver.config,
|
|
112
100
|
driver=await connect(cast(NucliaDBConnection, sync_driver.config)),
|
|
113
101
|
manager=await manager_connect(cast(NucliaDBConnection, sync_driver.config)),
|
|
114
102
|
_synonyms=None,
|
|
115
103
|
)
|
|
104
|
+
for sync_config_id in sync_driver.config.connection_ids:
|
|
105
|
+
await instance.resolve_sync_config(sync_config_id)
|
|
106
|
+
return instance
|
|
107
|
+
|
|
108
|
+
async def resolve_sync_config(
|
|
109
|
+
self, sync_config_id: str
|
|
110
|
+
) -> ExternalConnectionOutput:
|
|
111
|
+
if (
|
|
112
|
+
self.config.connection_ids
|
|
113
|
+
and sync_config_id not in self.config.connection_ids
|
|
114
|
+
):
|
|
115
|
+
raise ValueError(
|
|
116
|
+
f"Sync configuration {sync_config_id} is not configured for this driver"
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
existing_connection_ids = self.sync_configs.get(sync_config_id, [])
|
|
120
|
+
if existing_connection_ids:
|
|
121
|
+
return self.information[existing_connection_ids[0]]
|
|
122
|
+
|
|
123
|
+
UUID(sync_config_id)
|
|
124
|
+
response = await self.async_driver.get(f"/sync_config/{sync_config_id}")
|
|
125
|
+
response.raise_for_status()
|
|
126
|
+
connection_id = response.json()["external_connection"]["id"]
|
|
127
|
+
|
|
128
|
+
response = await self.async_driver.get(f"/external_connection/{connection_id}")
|
|
129
|
+
response.raise_for_status()
|
|
130
|
+
information = ExternalConnectionOutput(**response.json())
|
|
131
|
+
self.sync_configs[sync_config_id] = [connection_id]
|
|
132
|
+
self.information[connection_id] = information
|
|
133
|
+
return information
|
|
116
134
|
|
|
117
135
|
async def get_oauth_url(
|
|
118
136
|
self,
|