Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions py/core/main/api/v3/documents_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
WrappedRelationshipsResponse,
)
from core.utils import update_settings_from_dict
from core.utils.observability import create_r2r_trace_id
from shared.abstractions import IngestionMode

from ...abstractions import R2RProviders, R2RServices
Expand Down Expand Up @@ -398,6 +399,7 @@ async def create_document(

# Prepare workflow input
workflow_input = {
"langfuse_trace_id": create_r2r_trace_id(),
"document_id": str(document_id),
"chunks": [
chunk.model_dump(mode="json")
Expand Down Expand Up @@ -508,6 +510,7 @@ async def create_document(
)

workflow_input = {
"langfuse_trace_id": create_r2r_trace_id(),
"file_data": file_data,
"document_id": str(document_id),
"collection_ids": (
Expand Down
36 changes: 27 additions & 9 deletions py/core/main/orchestration/hatchet/ingestion_workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,12 @@
num_tokens,
update_settings_from_dict,
)
from core.utils.observability import (
build_r2r_ingestion_trace_context,
r2r_ingestion_trace_context,
reset_r2r_trace_context,
set_r2r_trace_context,
)

from ...services import IngestionService, IngestionServiceAdapter

Expand Down Expand Up @@ -59,9 +65,15 @@ def concurrency(self, context: Context) -> str:

@orchestration_provider.step(retries=0, timeout="60m")
async def parse(self, context: Context) -> dict:
input_data = context.workflow_input()["request"]
trace_token = set_r2r_trace_context(
build_r2r_ingestion_trace_context(
input_data,
task_id=str(context.workflow_run_id()),
)
)
try:
logger.info("Initiating ingestion workflow, step: parse")
input_data = context.workflow_input()["request"]
parsed_data = IngestionServiceAdapter.parse_ingest_file_input(
input_data
)
Expand Down Expand Up @@ -291,6 +303,8 @@ async def parse(self, context: Context) -> dict:
status_code=500,
detail=f"Error during ingestion: {str(e)}",
) from e
finally:
reset_r2r_trace_context(trace_token)

@orchestration_provider.failure()
async def on_failure(self, context: Context) -> None:
Expand Down Expand Up @@ -390,14 +404,18 @@ async def embed(self, context: Context) -> dict:
document_info = DocumentResponse(**document_info_dict)

extractions = context.step_output("ingest")["extractions"]

embedding_generator = self.ingestion_service.embed_document(
extractions
)
embeddings = [
embedding.model_dump()
async for embedding in embedding_generator
]
input_data = context.workflow_input()["request"]
with r2r_ingestion_trace_context(
input_data,
task_id=str(context.workflow_run_id()),
):
embedding_generator = self.ingestion_service.embed_document(
extractions
)
embeddings = [
embedding.model_dump()
async for embedding in embedding_generator
]

await self.ingestion_service.update_document_status(
document_info, status=IngestionStatus.STORING
Expand Down
15 changes: 15 additions & 0 deletions py/core/main/orchestration/simple/ingestion_workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,11 @@
num_tokens,
update_settings_from_dict,
)
from core.utils.observability import (
build_r2r_ingestion_trace_context,
reset_r2r_trace_context,
set_r2r_trace_context,
)

from ...services import IngestionService

Expand All @@ -24,6 +29,9 @@

def simple_ingestion_factory(service: IngestionService):
async def ingest_files(input_data):
trace_token = set_r2r_trace_context(
build_r2r_ingestion_trace_context(input_data)
)
document_info = None
try:
from core.base import IngestionStatus
Expand Down Expand Up @@ -209,6 +217,8 @@ async def ingest_files(input_data):
raise HTTPException(
status_code=500, detail=f"Error during ingestion: {str(e)}"
) from e
finally:
reset_r2r_trace_context(trace_token)

async def _ensure_collections_exists(
service: IngestionService,
Expand Down Expand Up @@ -284,6 +294,9 @@ async def _ensure_collections_exists(
raise e

async def ingest_chunks(input_data):
trace_token = set_r2r_trace_context(
build_r2r_ingestion_trace_context(input_data)
)
document_info = None
try:
from core.base import IngestionStatus
Expand Down Expand Up @@ -431,6 +444,8 @@ async def ingest_chunks(input_data):
status_code=500,
detail=f"Error during chunk ingestion: {str(e)}",
) from e
finally:
reset_r2r_trace_context(trace_token)

async def update_chunk(input_data):
from core.main import IngestionServiceAdapter
Expand Down
113 changes: 70 additions & 43 deletions py/core/main/services/ingestion_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
VectorTableName,
)
from core.base.api.models import User
from core.utils.observability import r2r_observation_context
from shared.abstractions import PDFParsingError, PopplerNotFoundError

from ..abstractions import R2RProviders
Expand Down Expand Up @@ -316,22 +317,24 @@ async def augment_document_info(
},
)

response = await self.providers.llm.aget_completion(
messages=messages,
generation_config=GenerationConfig(
model=self.config.ingestion.document_summary_model
or self.config.app.fast_llm
),
)
with r2r_observation_context("R2R: Document summary"):
response = await self.providers.llm.aget_completion(
messages=messages,
generation_config=GenerationConfig(
model=self.config.ingestion.document_summary_model
or self.config.app.fast_llm
),
)

document_info.summary = response.choices[0].message.content # type: ignore

if not document_info.summary:
raise ValueError("Expected a generated response.")

embedding = await self.providers.embedding.async_get_embedding(
text=document_info.summary,
)
with r2r_observation_context("R2R: Document summary embedding"):
embedding = await self.providers.embedding.async_get_embedding(
text=document_info.summary,
)
document_info.summary_embedding = embedding
return

Expand Down Expand Up @@ -366,9 +369,16 @@ async def process_batch(
for ex in batch
]
# Retrieve embeddings in bulk
vectors = await self.providers.embedding.async_get_embeddings(
texts, # list of strings
)
with r2r_observation_context(
"R2R: Chunk embedding",
metadata={
"r2r_chunk_count": len(batch),
"r2r_chunk_ids": [str(chunk.id) for chunk in batch],
},
):
vectors = await self.providers.embedding.async_get_embeddings(
texts, # list of strings
)
# Zip them back together
results = []
for raw_vector, extraction in zip(vectors, batch, strict=False):
Expand Down Expand Up @@ -733,35 +743,45 @@ async def _get_enriched_chunk_text(
]
try:
# Obtain the updated text from the LLM
updated_chunk_text = (
(
await self.providers.llm.aget_completion(
messages=await self.providers.database.prompts_handler.get_message_payload(
task_prompt_name=chunk_enrichment_settings.chunk_enrichment_prompt,
task_inputs={
"document_summary": document_summary or "None",
"chunk": chunk["text"],
"preceding_chunks": (
"\n".join(preceding_chunks)
if preceding_chunks
else "None"
),
"succeeding_chunks": (
"\n".join(succeeding_chunks)
if succeeding_chunks
else "None"
),
"chunk_size": self.config.ingestion.chunk_size
or 1024,
},
),
generation_config=chunk_enrichment_settings.generation_config
or GenerationConfig(model=self.config.app.fast_llm),
with r2r_observation_context(
"R2R: Chunk enrichment",
metadata={
"r2r_chunk_id": str(chunk["id"]),
"r2r_chunk_index": chunk_idx,
},
):
updated_chunk_text = (
(
await self.providers.llm.aget_completion(
messages=await self.providers.database.prompts_handler.get_message_payload(
task_prompt_name=chunk_enrichment_settings.chunk_enrichment_prompt,
task_inputs={
"document_summary": document_summary
or "None",
"chunk": chunk["text"],
"preceding_chunks": (
"\n".join(preceding_chunks)
if preceding_chunks
else "None"
),
"succeeding_chunks": (
"\n".join(succeeding_chunks)
if succeeding_chunks
else "None"
),
"chunk_size": self.config.ingestion.chunk_size
or 1024,
},
),
generation_config=chunk_enrichment_settings.generation_config
or GenerationConfig(
model=self.config.app.fast_llm
),
)
)
.choices[0]
.message.content
)
.choices[0]
.message.content
)
except Exception:
updated_chunk_text = chunk["text"]
chunk["metadata"]["chunk_enrichment_status"] = "failed"
Expand All @@ -775,9 +795,16 @@ async def _get_enriched_chunk_text(
chunk["metadata"]["chunk_enrichment_status"] = "failed"

# Re-embed
data = await self.providers.embedding.async_get_embedding(
updated_chunk_text
)
with r2r_observation_context(
"R2R: Enriched chunk embedding",
metadata={
"r2r_chunk_id": str(chunk["id"]),
"r2r_chunk_index": chunk_idx,
},
):
data = await self.providers.embedding.async_get_embedding(
updated_chunk_text
)
chunk["metadata"]["original_text"] = chunk["text"]

return VectorEntry(
Expand Down
Loading
Loading