hyperforge-nucliadb-agentic 1.0.0.post64__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.
- hyperforge_nucliadb_agentic/__init__.py +3 -0
- hyperforge_nucliadb_agentic/agent.py +1642 -0
- hyperforge_nucliadb_agentic/ask/__init__.py +5 -0
- hyperforge_nucliadb_agentic/ask/audit.py +439 -0
- hyperforge_nucliadb_agentic/ask/exceptions.py +50 -0
- hyperforge_nucliadb_agentic/ask/lifespan.py +21 -0
- hyperforge_nucliadb_agentic/ask/model.py +1299 -0
- hyperforge_nucliadb_agentic/ask/predict.py +431 -0
- hyperforge_nucliadb_agentic/ask/predict_models.py +78 -0
- hyperforge_nucliadb_agentic/ask/search/__init__.py +0 -0
- hyperforge_nucliadb_agentic/ask/search/ask.py +1182 -0
- hyperforge_nucliadb_agentic/ask/search/graph_strategy.py +1138 -0
- hyperforge_nucliadb_agentic/ask/search/highlight.py +93 -0
- hyperforge_nucliadb_agentic/ask/search/hydrator.py +29 -0
- hyperforge_nucliadb_agentic/ask/search/metrics.py +112 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/__init__.py +0 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/ask.py +70 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/fetcher.py +192 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/find.py +729 -0
- hyperforge_nucliadb_agentic/ask/search/prompt.py +1298 -0
- hyperforge_nucliadb_agentic/ask/search/rank_fusion.py +157 -0
- hyperforge_nucliadb_agentic/ask/search/rerankers.py +161 -0
- hyperforge_nucliadb_agentic/ask/search/retrieval.py +750 -0
- hyperforge_nucliadb_agentic/ask/search/rpc.py +192 -0
- hyperforge_nucliadb_agentic/ask/settings.py +9 -0
- hyperforge_nucliadb_agentic/ask/utils/ids.py +188 -0
- hyperforge_nucliadb_agentic/ask/utils/proto.py +6 -0
- hyperforge_nucliadb_agentic/ask/utils/responses.py +6 -0
- hyperforge_nucliadb_agentic/ask/utils/text_blocks.py +51 -0
- hyperforge_nucliadb_agentic/config.py +47 -0
- hyperforge_nucliadb_agentic/internal_driver.py +87 -0
- hyperforge_nucliadb_agentic/py.typed +0 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/METADATA +24 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/RECORD +35 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,439 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import contextvars
|
|
3
|
+
import time
|
|
4
|
+
from datetime import datetime, timezone
|
|
5
|
+
from typing import cast
|
|
6
|
+
|
|
7
|
+
import backoff
|
|
8
|
+
import mmh3
|
|
9
|
+
import nats
|
|
10
|
+
from fastapi import Request
|
|
11
|
+
from hyperforge.feature_flag import Features, has_feature
|
|
12
|
+
from nucliadb_models.retrieval import RawQuery, RetrievalRequest
|
|
13
|
+
from nucliadb_models.search import (
|
|
14
|
+
NucliaDBClientType,
|
|
15
|
+
)
|
|
16
|
+
from nucliadb_protos import audit_pb2, utils_pb2
|
|
17
|
+
from nucliadb_protos.audit_pb2 import (
|
|
18
|
+
AuditRequest,
|
|
19
|
+
ChatContext,
|
|
20
|
+
RetrievedContext,
|
|
21
|
+
)
|
|
22
|
+
from nucliadb_telemetry.jetstream import get_traced_jetstream, get_traced_nats_client
|
|
23
|
+
from nucliadb_utils import logger
|
|
24
|
+
from nucliadb_utils.settings import AuditSettings
|
|
25
|
+
from nucliadb_utils.utilities import Utility, clean_utility, get_utility, set_utility
|
|
26
|
+
from opentelemetry.trace import INVALID_SPAN, format_trace_id, get_current_span
|
|
27
|
+
from starlette.background import BackgroundTask
|
|
28
|
+
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
|
|
29
|
+
from starlette.responses import Response
|
|
30
|
+
|
|
31
|
+
from hyperforge_nucliadb_agentic.ask.model import (
|
|
32
|
+
AskRequest,
|
|
33
|
+
ChatContextMessage,
|
|
34
|
+
PromptContext,
|
|
35
|
+
PromptContextOrder,
|
|
36
|
+
)
|
|
37
|
+
from hyperforge_nucliadb_agentic.ask.predict import AnswerStatusCode
|
|
38
|
+
from hyperforge_nucliadb_agentic.ask.utils.proto import client_type
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class RequestContext:
|
|
42
|
+
def __init__(self: "RequestContext"):
|
|
43
|
+
self.audit_request: AuditRequest = AuditRequest()
|
|
44
|
+
self.start_time: float = time.monotonic()
|
|
45
|
+
self.path: str = ""
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
request_context_var = contextvars.ContextVar[RequestContext | None](
|
|
49
|
+
"request_context", default=None
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def get_trace_id() -> str | None:
|
|
54
|
+
span = get_current_span()
|
|
55
|
+
if span is INVALID_SPAN:
|
|
56
|
+
return None
|
|
57
|
+
return format_trace_id(span.get_span_context().trace_id)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def get_request_context() -> RequestContext | None:
|
|
61
|
+
return request_context_var.get()
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class AuditMiddleware(BaseHTTPMiddleware):
|
|
65
|
+
async def dispatch(
|
|
66
|
+
self, request: Request, call_next: RequestResponseEndpoint
|
|
67
|
+
) -> Response:
|
|
68
|
+
context = RequestContext()
|
|
69
|
+
token = request_context_var.set(context)
|
|
70
|
+
context.audit_request.time.FromDatetime(datetime.now(tz=timezone.utc))
|
|
71
|
+
context.audit_request.trace_id = get_trace_id() or ""
|
|
72
|
+
context.path = request.url.path
|
|
73
|
+
|
|
74
|
+
response = await call_next(request)
|
|
75
|
+
|
|
76
|
+
# This task will run when the response finishes streaming
|
|
77
|
+
# When dealing with streaming responses, AND if we depend on any state that only will be available once
|
|
78
|
+
# the request is fully finished, the response we have after the dispatch call_next is not enough, as
|
|
79
|
+
# there, no iteration of the streaming response has been done yet.
|
|
80
|
+
response.background = BackgroundTask(self.enqueue_pending, context)
|
|
81
|
+
|
|
82
|
+
# It is safe to reset the context here since the asyncio task for generating the streaming response is
|
|
83
|
+
# already running. If we want to spawn a different task during streaming and we want that task be able
|
|
84
|
+
# to read the context_var, we need to manually pass the context into that task.
|
|
85
|
+
request_context_var.reset(token)
|
|
86
|
+
|
|
87
|
+
return response
|
|
88
|
+
|
|
89
|
+
def enqueue_pending(self, context: RequestContext):
|
|
90
|
+
if context.audit_request.kbid:
|
|
91
|
+
# an audit request with no kbid makes no sense, we use this as an heuristic
|
|
92
|
+
# mark that no audit has been set during this request
|
|
93
|
+
|
|
94
|
+
context.audit_request.request_time = time.monotonic() - context.start_time
|
|
95
|
+
audit = get_audit()
|
|
96
|
+
if audit is not None:
|
|
97
|
+
audit.send(context.audit_request)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
class StreamAuditStorage:
|
|
101
|
+
task: asyncio.Task | None
|
|
102
|
+
initialized: bool
|
|
103
|
+
queue: asyncio.Queue
|
|
104
|
+
|
|
105
|
+
def __init__(
|
|
106
|
+
self,
|
|
107
|
+
nats_servers: list[str],
|
|
108
|
+
nats_target: str,
|
|
109
|
+
partitions: int,
|
|
110
|
+
seed: int,
|
|
111
|
+
nats_creds: str | None,
|
|
112
|
+
service: str,
|
|
113
|
+
):
|
|
114
|
+
self.nats_servers = nats_servers
|
|
115
|
+
self.nats_creds = nats_creds
|
|
116
|
+
self.nats_target = nats_target
|
|
117
|
+
self.partitions = partitions
|
|
118
|
+
self.seed = seed
|
|
119
|
+
self.queue = asyncio.Queue()
|
|
120
|
+
self.service = service
|
|
121
|
+
self.task = None
|
|
122
|
+
self.initialized = False
|
|
123
|
+
|
|
124
|
+
def get_partition(self, kbid: str):
|
|
125
|
+
return mmh3.hash(kbid, self.seed, signed=False) % self.partitions
|
|
126
|
+
|
|
127
|
+
async def disconnected_cb(self):
|
|
128
|
+
logger.info("Got disconnected from NATS!")
|
|
129
|
+
|
|
130
|
+
async def reconnected_cb(self):
|
|
131
|
+
# See who we are connected to on reconnect.
|
|
132
|
+
logger.info(f"Got reconnected to NATS {self.nc.connected_url}") # type: ignore
|
|
133
|
+
|
|
134
|
+
async def error_cb(self, e):
|
|
135
|
+
logger.error(f"There was an error connecting to NATS audit: {e}", exc_info=True)
|
|
136
|
+
|
|
137
|
+
async def closed_cb(self):
|
|
138
|
+
logger.info("Connection is closed on NATS")
|
|
139
|
+
|
|
140
|
+
async def initialize(self):
|
|
141
|
+
options = {
|
|
142
|
+
"error_cb": self.error_cb,
|
|
143
|
+
"closed_cb": self.closed_cb,
|
|
144
|
+
"reconnected_cb": self.reconnected_cb,
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
if self.nats_creds:
|
|
148
|
+
options["user_credentials"] = self.nats_creds # type: ignore
|
|
149
|
+
|
|
150
|
+
if len(self.nats_servers) > 0:
|
|
151
|
+
options["servers"] = self.nats_servers # type: ignore
|
|
152
|
+
|
|
153
|
+
nc = await nats.connect(**options) # type: ignore
|
|
154
|
+
self.nc = get_traced_nats_client(nc, self.service)
|
|
155
|
+
|
|
156
|
+
self.js = get_traced_jetstream(self.nc, self.service)
|
|
157
|
+
self.task = asyncio.create_task(self.run())
|
|
158
|
+
|
|
159
|
+
self.initialized = True
|
|
160
|
+
|
|
161
|
+
async def finalize(self):
|
|
162
|
+
if self.task is not None:
|
|
163
|
+
self.task.cancel()
|
|
164
|
+
if self.nc:
|
|
165
|
+
await self.nc.flush()
|
|
166
|
+
await self.nc.close()
|
|
167
|
+
self.nc = None
|
|
168
|
+
|
|
169
|
+
async def run(self):
|
|
170
|
+
while True:
|
|
171
|
+
item_dequeued = False
|
|
172
|
+
try:
|
|
173
|
+
audit = await self.queue.get()
|
|
174
|
+
item_dequeued = True
|
|
175
|
+
await self._send(audit)
|
|
176
|
+
except (asyncio.CancelledError, KeyboardInterrupt, RuntimeError):
|
|
177
|
+
return
|
|
178
|
+
except Exception: # pragma: no cover
|
|
179
|
+
logger.exception("Could not send audit", stack_info=True)
|
|
180
|
+
finally:
|
|
181
|
+
if item_dequeued:
|
|
182
|
+
self.queue.task_done()
|
|
183
|
+
|
|
184
|
+
def send(self, message: AuditRequest):
|
|
185
|
+
self.queue.put_nowait(message)
|
|
186
|
+
|
|
187
|
+
@backoff.on_exception(
|
|
188
|
+
backoff.expo, (Exception,), jitter=backoff.random_jitter, max_tries=4
|
|
189
|
+
)
|
|
190
|
+
async def _send(self, message: AuditRequest):
|
|
191
|
+
if self.js is None: # pragma: no cover
|
|
192
|
+
raise AttributeError()
|
|
193
|
+
|
|
194
|
+
partition = self.get_partition(message.kbid)
|
|
195
|
+
|
|
196
|
+
res = await self.js.publish(
|
|
197
|
+
self.nats_target.format(partition=partition, type=message.type),
|
|
198
|
+
message.SerializeToString(),
|
|
199
|
+
)
|
|
200
|
+
logger.debug(
|
|
201
|
+
f"Pushed message to audit. kb: {message.kbid}, resource: {message.rid}, partition: {partition}"
|
|
202
|
+
)
|
|
203
|
+
return res.seq
|
|
204
|
+
|
|
205
|
+
def retrieve(
|
|
206
|
+
self,
|
|
207
|
+
retrieval_time: float,
|
|
208
|
+
resources: int,
|
|
209
|
+
retrieval_request: RetrievalRequest,
|
|
210
|
+
):
|
|
211
|
+
context = get_request_context()
|
|
212
|
+
if context is None:
|
|
213
|
+
return
|
|
214
|
+
|
|
215
|
+
auditrequest = context.audit_request
|
|
216
|
+
|
|
217
|
+
auditrequest.retrieval_time = retrieval_time
|
|
218
|
+
auditrequest.resources = resources
|
|
219
|
+
|
|
220
|
+
auditrequest.search.result_per_page = retrieval_request.top_k
|
|
221
|
+
|
|
222
|
+
if (
|
|
223
|
+
isinstance(retrieval_request.query, RawQuery)
|
|
224
|
+
and retrieval_request.query.keyword is not None
|
|
225
|
+
):
|
|
226
|
+
auditrequest.search.body = retrieval_request.query.keyword.query
|
|
227
|
+
auditrequest.search.min_score_bm25 = (
|
|
228
|
+
retrieval_request.query.keyword.min_score
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
if (
|
|
232
|
+
isinstance(retrieval_request.query, RawQuery)
|
|
233
|
+
and retrieval_request.query.semantic is not None
|
|
234
|
+
):
|
|
235
|
+
auditrequest.search.vector.extend(retrieval_request.query.semantic.query)
|
|
236
|
+
auditrequest.search.min_score_bm25 = (
|
|
237
|
+
retrieval_request.query.semantic.min_score
|
|
238
|
+
)
|
|
239
|
+
auditrequest.search.vectorset = retrieval_request.query.semantic.vectorset
|
|
240
|
+
|
|
241
|
+
if retrieval_request.filters.filter_expression is not None:
|
|
242
|
+
# NOTE: this filter is a dump of the API models. NucliaDB
|
|
243
|
+
# implementation uses the filter expression proto in JSON format.
|
|
244
|
+
# However, neither we have the proto nor we want to do a costly
|
|
245
|
+
# conversion (we'd need to query nucliadb for a slug to rid
|
|
246
|
+
# conversion in order to build the proto)
|
|
247
|
+
auditrequest.search.filter = retrieval_request.filters.model_dump_json()
|
|
248
|
+
|
|
249
|
+
if retrieval_request.filters.security is not None:
|
|
250
|
+
security_pb = utils_pb2.Security()
|
|
251
|
+
for group_id in retrieval_request.filters.security.groups:
|
|
252
|
+
if group_id not in security_pb.access_groups:
|
|
253
|
+
security_pb.access_groups.append(group_id)
|
|
254
|
+
auditrequest.search.security.CopyFrom(security_pb)
|
|
255
|
+
|
|
256
|
+
def ask(
|
|
257
|
+
self,
|
|
258
|
+
kbid: str,
|
|
259
|
+
user: str,
|
|
260
|
+
client_type: int,
|
|
261
|
+
origin: str,
|
|
262
|
+
ask_request: AskRequest,
|
|
263
|
+
question: str,
|
|
264
|
+
rephrased_question: str | None,
|
|
265
|
+
retrieval_rephrased_question: str | None,
|
|
266
|
+
chat_context: list[ChatContext],
|
|
267
|
+
retrieved_context: list[RetrievedContext],
|
|
268
|
+
answer: str | None,
|
|
269
|
+
reasoning: str | None,
|
|
270
|
+
learning_id: str | None,
|
|
271
|
+
status_code: int,
|
|
272
|
+
model: str | None,
|
|
273
|
+
rephrase_time: float | None = None,
|
|
274
|
+
generative_answer_time: float | None = None,
|
|
275
|
+
generative_answer_first_chunk_time: float | None = None,
|
|
276
|
+
generative_reasoning_first_chunk_time: float | None = None,
|
|
277
|
+
):
|
|
278
|
+
if not has_feature(Features.AUDIT_RAO_ASK_ENDPOINT):
|
|
279
|
+
return
|
|
280
|
+
|
|
281
|
+
rcontext = get_request_context()
|
|
282
|
+
if rcontext is None:
|
|
283
|
+
return
|
|
284
|
+
|
|
285
|
+
audit_request = rcontext.audit_request
|
|
286
|
+
|
|
287
|
+
audit_request.type = AuditRequest.AuditType.ASK
|
|
288
|
+
audit_request.origin = origin
|
|
289
|
+
audit_request.client_type = client_type # type: ignore
|
|
290
|
+
audit_request.userid = user
|
|
291
|
+
audit_request.kbid = kbid
|
|
292
|
+
audit_request.user_request = ask_request.model_dump_json(exclude_unset=True)
|
|
293
|
+
if rephrase_time is not None:
|
|
294
|
+
audit_request.rephrase_time = rephrase_time
|
|
295
|
+
if generative_answer_time is not None:
|
|
296
|
+
audit_request.generative_answer_time = generative_answer_time
|
|
297
|
+
if generative_answer_first_chunk_time is not None:
|
|
298
|
+
audit_request.generative_answer_first_chunk_time = (
|
|
299
|
+
generative_answer_first_chunk_time
|
|
300
|
+
)
|
|
301
|
+
if generative_reasoning_first_chunk_time is not None:
|
|
302
|
+
audit_request.generative_reasoning_first_chunk_time = (
|
|
303
|
+
generative_reasoning_first_chunk_time
|
|
304
|
+
)
|
|
305
|
+
|
|
306
|
+
if retrieval_rephrased_question is not None:
|
|
307
|
+
audit_request.retrieval_rephrased_question = retrieval_rephrased_question
|
|
308
|
+
|
|
309
|
+
audit_request.chat.question = question
|
|
310
|
+
audit_request.chat.chat_context.extend(chat_context)
|
|
311
|
+
audit_request.chat.retrieved_context.extend(retrieved_context)
|
|
312
|
+
if learning_id is not None:
|
|
313
|
+
audit_request.chat.learning_id = learning_id
|
|
314
|
+
if rephrased_question is not None:
|
|
315
|
+
audit_request.chat.rephrased_question = rephrased_question
|
|
316
|
+
if answer is not None:
|
|
317
|
+
audit_request.chat.answer = answer
|
|
318
|
+
if reasoning is not None:
|
|
319
|
+
audit_request.chat.reasoning = reasoning
|
|
320
|
+
|
|
321
|
+
audit_request.chat.status_code = status_code
|
|
322
|
+
if model is not None:
|
|
323
|
+
audit_request.chat.model = model
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
def get_audit() -> StreamAuditStorage | None:
|
|
327
|
+
return get_utility(Utility.AUDIT)
|
|
328
|
+
|
|
329
|
+
|
|
330
|
+
async def start_audit_utility(
|
|
331
|
+
service: str, audit_settings: AuditSettings
|
|
332
|
+
) -> StreamAuditStorage:
|
|
333
|
+
audit_utility = StreamAuditStorage(
|
|
334
|
+
nats_creds=audit_settings.audit_jetstream_auth,
|
|
335
|
+
nats_servers=audit_settings.audit_jetstream_servers,
|
|
336
|
+
nats_target=cast(str, audit_settings.audit_jetstream_target),
|
|
337
|
+
partitions=audit_settings.audit_partitions,
|
|
338
|
+
seed=audit_settings.audit_hash_seed,
|
|
339
|
+
service=service,
|
|
340
|
+
)
|
|
341
|
+
await audit_utility.initialize()
|
|
342
|
+
set_utility(Utility.AUDIT, audit_utility)
|
|
343
|
+
return audit_utility
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
async def stop_audit_utility():
|
|
347
|
+
audit_utility = get_utility(Utility.AUDIT)
|
|
348
|
+
if audit_utility is None:
|
|
349
|
+
return
|
|
350
|
+
clean_utility(Utility.AUDIT)
|
|
351
|
+
await audit_utility.finalize()
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
class ChatAuditor:
|
|
355
|
+
def __init__(
|
|
356
|
+
self,
|
|
357
|
+
kbid: str,
|
|
358
|
+
user_id: str,
|
|
359
|
+
client_type: NucliaDBClientType,
|
|
360
|
+
origin: str,
|
|
361
|
+
ask_request: AskRequest,
|
|
362
|
+
user_query: str,
|
|
363
|
+
rephrased_query: str | None,
|
|
364
|
+
retrieval_rephrased_query: str | None,
|
|
365
|
+
chat_history: list[ChatContextMessage],
|
|
366
|
+
learning_id: str | None,
|
|
367
|
+
query_context: PromptContext,
|
|
368
|
+
query_context_order: PromptContextOrder,
|
|
369
|
+
model: str | None,
|
|
370
|
+
):
|
|
371
|
+
self.kbid = kbid
|
|
372
|
+
self.user_id = user_id
|
|
373
|
+
self.client_type = client_type
|
|
374
|
+
self.origin = origin
|
|
375
|
+
self.ask_request = ask_request
|
|
376
|
+
self.user_query = user_query
|
|
377
|
+
self.rephrased_query = rephrased_query
|
|
378
|
+
self.retrieval_rephrased_query = retrieval_rephrased_query
|
|
379
|
+
self.chat_history = chat_history
|
|
380
|
+
self.learning_id = learning_id
|
|
381
|
+
self.query_context = query_context
|
|
382
|
+
self.query_context_order = query_context_order
|
|
383
|
+
self.model = model
|
|
384
|
+
|
|
385
|
+
def audit(
|
|
386
|
+
self,
|
|
387
|
+
text_answer: bytes,
|
|
388
|
+
text_reasoning: str | None,
|
|
389
|
+
generative_answer_time: float,
|
|
390
|
+
generative_answer_first_chunk_time: float,
|
|
391
|
+
generative_reasoning_first_chunk_time: float | None,
|
|
392
|
+
rephrase_time: float | None,
|
|
393
|
+
status_code: AnswerStatusCode,
|
|
394
|
+
):
|
|
395
|
+
audit = get_audit()
|
|
396
|
+
if audit is None:
|
|
397
|
+
return
|
|
398
|
+
|
|
399
|
+
if (
|
|
400
|
+
status_code == AnswerStatusCode.NO_CONTEXT
|
|
401
|
+
or status_code == AnswerStatusCode.NO_RETRIEVAL_DATA
|
|
402
|
+
): # We don't want to audit "Not enough context to answer this." and instead set a None.
|
|
403
|
+
audit_answer = None
|
|
404
|
+
else:
|
|
405
|
+
audit_answer = text_answer.decode()
|
|
406
|
+
|
|
407
|
+
# Append chat history
|
|
408
|
+
chat_history_context = [
|
|
409
|
+
audit_pb2.ChatContext(author=message.author, text=message.text)
|
|
410
|
+
for message in self.chat_history
|
|
411
|
+
]
|
|
412
|
+
|
|
413
|
+
# Append paragraphs retrieved on this chat
|
|
414
|
+
chat_retrieved_context = [
|
|
415
|
+
audit_pb2.RetrievedContext(text_block_id=paragraph_id, text=text)
|
|
416
|
+
for paragraph_id, text in self.query_context.items()
|
|
417
|
+
]
|
|
418
|
+
|
|
419
|
+
audit.ask(
|
|
420
|
+
self.kbid,
|
|
421
|
+
self.user_id,
|
|
422
|
+
client_type(self.client_type),
|
|
423
|
+
self.origin,
|
|
424
|
+
ask_request=self.ask_request,
|
|
425
|
+
question=self.user_query,
|
|
426
|
+
generative_answer_time=generative_answer_time,
|
|
427
|
+
generative_answer_first_chunk_time=generative_answer_first_chunk_time,
|
|
428
|
+
generative_reasoning_first_chunk_time=generative_reasoning_first_chunk_time,
|
|
429
|
+
rephrase_time=rephrase_time,
|
|
430
|
+
rephrased_question=self.rephrased_query,
|
|
431
|
+
retrieval_rephrased_question=self.retrieval_rephrased_query,
|
|
432
|
+
chat_context=chat_history_context,
|
|
433
|
+
retrieved_context=chat_retrieved_context,
|
|
434
|
+
answer=audit_answer,
|
|
435
|
+
reasoning=text_reasoning,
|
|
436
|
+
learning_id=self.learning_id,
|
|
437
|
+
status_code=int(status_code.value),
|
|
438
|
+
model=self.model,
|
|
439
|
+
)
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
"""Package exceptions"""
|
|
2
|
+
|
|
3
|
+
from nucliadb_models.search import KnowledgeboxFindResults, PreQueryResult
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class KnowledgeBoxNotFound(Exception):
|
|
7
|
+
pass
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class ResourceNotFoundError(Exception):
|
|
11
|
+
pass
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class InvalidQueryError(Exception):
|
|
15
|
+
"""Raised when parsing a query containing an invalid parameter"""
|
|
16
|
+
|
|
17
|
+
def __init__(self, param: str, reason: str):
|
|
18
|
+
self.param = param
|
|
19
|
+
self.reason = reason
|
|
20
|
+
super().__init__(f"Invalid query. Error in {param}: {reason}")
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class InternalParserError(ValueError):
|
|
24
|
+
"""Raised when parsing fails due to some internal error"""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class NoRetrievalResultsError(Exception):
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
main: KnowledgeboxFindResults | None = None,
|
|
31
|
+
prequeries: list[PreQueryResult] | None = None,
|
|
32
|
+
prefilters: list[PreQueryResult] | None = None,
|
|
33
|
+
):
|
|
34
|
+
self.main_query = main
|
|
35
|
+
self.prequeries = prequeries
|
|
36
|
+
self.prefilters = prefilters
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class AnswerJsonSchemaTooLong(Exception):
|
|
40
|
+
pass
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class IncompleteFindResultsError(Exception):
|
|
44
|
+
pass
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class NucliaDBError(Exception):
|
|
48
|
+
"""NucliaDB internal error raised on 500 and other NucliaDB errors"""
|
|
49
|
+
|
|
50
|
+
pass
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from contextlib import asynccontextmanager
|
|
2
|
+
|
|
3
|
+
from fastapi import FastAPI
|
|
4
|
+
from nucliadb_telemetry.utils import clean_telemetry, setup_telemetry
|
|
5
|
+
|
|
6
|
+
from hyperforge_nucliadb_agentic.ask import SERVICE_NAME
|
|
7
|
+
from hyperforge_nucliadb_agentic.ask.predict import (
|
|
8
|
+
start_predict_engine,
|
|
9
|
+
stop_predict_engine,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@asynccontextmanager
|
|
14
|
+
async def lifespan(app: FastAPI):
|
|
15
|
+
await setup_telemetry(SERVICE_NAME)
|
|
16
|
+
await start_predict_engine()
|
|
17
|
+
|
|
18
|
+
yield
|
|
19
|
+
|
|
20
|
+
await stop_predict_engine()
|
|
21
|
+
await clean_telemetry(SERVICE_NAME)
|