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.
Files changed (35) hide show
  1. hyperforge_nucliadb_agentic/__init__.py +3 -0
  2. hyperforge_nucliadb_agentic/agent.py +1642 -0
  3. hyperforge_nucliadb_agentic/ask/__init__.py +5 -0
  4. hyperforge_nucliadb_agentic/ask/audit.py +439 -0
  5. hyperforge_nucliadb_agentic/ask/exceptions.py +50 -0
  6. hyperforge_nucliadb_agentic/ask/lifespan.py +21 -0
  7. hyperforge_nucliadb_agentic/ask/model.py +1299 -0
  8. hyperforge_nucliadb_agentic/ask/predict.py +431 -0
  9. hyperforge_nucliadb_agentic/ask/predict_models.py +78 -0
  10. hyperforge_nucliadb_agentic/ask/search/__init__.py +0 -0
  11. hyperforge_nucliadb_agentic/ask/search/ask.py +1182 -0
  12. hyperforge_nucliadb_agentic/ask/search/graph_strategy.py +1138 -0
  13. hyperforge_nucliadb_agentic/ask/search/highlight.py +93 -0
  14. hyperforge_nucliadb_agentic/ask/search/hydrator.py +29 -0
  15. hyperforge_nucliadb_agentic/ask/search/metrics.py +112 -0
  16. hyperforge_nucliadb_agentic/ask/search/parsers/__init__.py +0 -0
  17. hyperforge_nucliadb_agentic/ask/search/parsers/ask.py +70 -0
  18. hyperforge_nucliadb_agentic/ask/search/parsers/fetcher.py +192 -0
  19. hyperforge_nucliadb_agentic/ask/search/parsers/find.py +729 -0
  20. hyperforge_nucliadb_agentic/ask/search/prompt.py +1298 -0
  21. hyperforge_nucliadb_agentic/ask/search/rank_fusion.py +157 -0
  22. hyperforge_nucliadb_agentic/ask/search/rerankers.py +161 -0
  23. hyperforge_nucliadb_agentic/ask/search/retrieval.py +750 -0
  24. hyperforge_nucliadb_agentic/ask/search/rpc.py +192 -0
  25. hyperforge_nucliadb_agentic/ask/settings.py +9 -0
  26. hyperforge_nucliadb_agentic/ask/utils/ids.py +188 -0
  27. hyperforge_nucliadb_agentic/ask/utils/proto.py +6 -0
  28. hyperforge_nucliadb_agentic/ask/utils/responses.py +6 -0
  29. hyperforge_nucliadb_agentic/ask/utils/text_blocks.py +51 -0
  30. hyperforge_nucliadb_agentic/config.py +47 -0
  31. hyperforge_nucliadb_agentic/internal_driver.py +87 -0
  32. hyperforge_nucliadb_agentic/py.typed +0 -0
  33. hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/METADATA +24 -0
  34. hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/RECORD +35 -0
  35. hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/WHEEL +4 -0
@@ -0,0 +1,5 @@
1
+ import logging
2
+
3
+ SERVICE_NAME = "nucliadb_agentic_agentic.ask"
4
+
5
+ logger = logging.getLogger(SERVICE_NAME)
@@ -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)