Skip to content
Merged
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
30 changes: 24 additions & 6 deletions apps/api/app/api/v1/routes/retrieval.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

from __future__ import annotations

from typing import Literal
from typing import Any, Literal

from app.api.dependencies.current_user import with_current_user
from app.services.rate_limit.data_structures import CurrentUser
Expand All @@ -11,6 +11,7 @@
from sqlalchemy.ext.asyncio import AsyncSession

from shared.core.database import get_db
from shared.models.schemas.llm_config import LLMConfig
from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace
from shared.services.retrieval.app_service import run_retrieval_query
from shared.services.retrieval.settings import DEFAULT_TOP_K, VALID_CHUNK_TYPES, normalize_chunk_types
Expand Down Expand Up @@ -129,12 +130,14 @@ class RetrievalQueryResponse(BaseModel):
)


@router.post("/query", response_model=RetrievalQueryResponse)
async def query_retrieval(
async def execute_retrieval_query(
payload: RetrievalQueryRequest,
current_user: CurrentUser = Depends(with_current_user),
db: AsyncSession = Depends(get_db),
):
current_user: CurrentUser,
db: AsyncSession,
*,
llm_config: LLMConfig | None = None,
) -> dict[str, Any]:
"""Shared retrieval execution used by v1 and v2 route handlers."""
# Resolve chunk_types: explicit field takes precedence over legacy data_type
if payload.chunk_types is not None:
resolved_chunk_types = normalize_chunk_types(payload.chunk_types)
Expand Down Expand Up @@ -166,4 +169,19 @@ async def query_retrieval(
threshold=payload.threshold,
internal_recall_k=payload.internal_recall_k,
use_agentic=payload.use_agentic,
llm_config=llm_config,
)


@router.post("/query", response_model=RetrievalQueryResponse)
async def query_retrieval(
payload: RetrievalQueryRequest,
current_user: CurrentUser = Depends(with_current_user),
db: AsyncSession = Depends(get_db),
):
return await execute_retrieval_query(
payload,
current_user,
db,
llm_config=None,
)
47 changes: 45 additions & 2 deletions apps/api/app/api/v2/routes/retrieval.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,48 @@
"""Retrieval API v2 routes."""

from app.api.v1.routes.retrieval import router
from __future__ import annotations

__all__ = ["router"]
from app.api.dependencies.current_user import with_current_user
from app.api.v1.routes.retrieval import (
RetrievalQueryRequest,
RetrievalQueryResponse,
execute_retrieval_query,
)
from app.services.rate_limit.data_structures import CurrentUser
from fastapi import APIRouter, Depends
from pydantic import Field
from sqlalchemy.ext.asyncio import AsyncSession

from shared.core.database import get_db
from shared.models.schemas.llm_config import LLMConfig

router = APIRouter(tags=["Retrieval"])


class RetrievalQueryRequestV2(RetrievalQueryRequest):
"""v2 retrieval query request with optional BYOK LLM credentials."""

llm_config: LLMConfig | None = Field(
None,
description=(
"Optional bring-your-own-key OpenAI-compatible LLM credentials. "
"Flat {api_key, model, base_url} applies to both channels. "
"Use models.{text,vision} for different model ids on the same "
"endpoint, or text/vision objects for different endpoints. "
"When omitted, server defaults are used."
),
)


@router.post("/query", response_model=RetrievalQueryResponse)
async def query_retrieval(
payload: RetrievalQueryRequestV2,
current_user: CurrentUser = Depends(with_current_user),
db: AsyncSession = Depends(get_db),
):
return await execute_retrieval_query(
payload,
current_user,
db,
llm_config=payload.llm_config,
)
5 changes: 3 additions & 2 deletions apps/worker/app/services/connect_builder/summary_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ def _llm_summarize(snippets_text: str, node_name: str, max_tokens: int = 100) ->
"""
try:
from shared.services.ai.prompt_service import build_prompt, _detect_text_language
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client

# Deterministic language lock — see prompt_service._language_directive
detected_lang = _detect_text_language(snippets_text)
Expand All @@ -58,7 +58,8 @@ def _llm_summarize(snippets_text: str, node_name: str, max_tokens: int = 100) ->
{"role": "system", "content": "you are a helpful assistant"},
{"role": "user", "content": prompt},
]
resp = get_openai_client().chat_completion(
client, _ = get_text_client()
resp = client.chat_completion(
messages=messages,
timeout=60,
max_tokens=max_tokens,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -235,9 +235,9 @@ def _next_decision(self, round_index: int) -> tuple[ReflexionDecision, ToolResul
)
start = time.monotonic()
try:
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client

client = get_openai_client(model=model)
client, model = get_text_client(requested_model=model)
raw, usage = client.chat_completion_with_usage(
messages=[{"role": "user", "content": prompt}],
model=model,
Expand Down
4 changes: 2 additions & 2 deletions apps/worker/app/services/document_agent/planner/planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -286,9 +286,9 @@ def propose(self) -> tuple[DocumentProfile, ReflexionDecision, ToolResult]:
logger.warning("[document_agent] planner png attach failed: {}", exc)

try:
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_vision_client

client = get_openai_client(model=model)
client, model = get_vision_client(requested_model=model)
raw, usage = client.chat_completion_with_usage(
messages=cast(Any, [{"role": "user", "content": content_parts}]),
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -328,9 +328,9 @@ def verify_section_page_choice(

start = time.monotonic()
try:
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_vision_client

client = get_openai_client(model=model)
client, model = get_vision_client(requested_model=model)
raw, usage = client.chat_completion_with_usage(
messages=cast(Any, [{"role": "user", "content": content_parts}]),
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -299,9 +299,9 @@ def decide_next_action(
logger.warning("[page_locate.subagent] planner budget exhausted for title={!r}", title)
return None
try:
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client

client = get_openai_client(model=model)
client, model = get_text_client(requested_model=model)
raw, usage = client.chat_completion_with_usage(
messages=[{"role": "user", "content": prompt}],
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@ def _vlm_confirm_anchors(
budget: Any | None = None,
) -> tuple[list[TocAnchorPage], bool, list[TocEvidence]]:
"""Phase 1: send all anchor PNGs to VLM, ask which are real TOC starts."""
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_vision_client

if not anchor_pages:
return [], False, []
Expand Down Expand Up @@ -125,7 +125,8 @@ def _vlm_confirm_anchors(
return [], True, []

try:
client = get_openai_client(model=model)
client, resolved_model = get_vision_client(requested_model=model)
model = resolved_model or model
raw, usage = client.chat_completion_with_usage(
messages=messages,
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -78,9 +78,9 @@ def inspect_pages(ctx: ToolContext, args: dict[str, Any]) -> ToolResult:
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_b64}"}}
)
try:
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_vision_client

client = get_openai_client(model=model)
client, model = get_vision_client(requested_model=model)
raw, usage = client.chat_completion_with_usage(
messages=cast(Any, [{"role": "user", "content": content_parts}]),
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -738,9 +738,9 @@ def propose_shard_plan(ctx: ToolContext, _args: dict[str, Any]) -> ToolResult:
if model and ctx.budget.try_reserve("plan", prompt_tokens_est):
try:
llm_attempted = True
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client

client = get_openai_client(model=model)
client, model = get_text_client(requested_model=model)
raw_response, usage = client.chat_completion_with_usage(
messages=[{"role": "user", "content": prompt}],
model=model,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -162,7 +162,7 @@ def vlm_extract_toc_batch(
BatchTocResult with per-page classification and extracted entries.
"""
from loguru import logger
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_vision_client

if not page_pngs:
return BatchTocResult(
Expand Down Expand Up @@ -190,7 +190,8 @@ def vlm_extract_toc_batch(
)

start = time.monotonic()
client = get_openai_client(model=model)
client, resolved_model = get_vision_client(requested_model=model)
model = resolved_model or model
raw, usage = client.chat_completion_with_usage(
messages=cast(Any, [{"role": "user", "content": content_parts}]),
model=model,
Expand Down
3 changes: 3 additions & 0 deletions apps/worker/app/services/document_ingestion/processing_run.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
)
from shared.core.exceptions.domain_exceptions import ValidationException
from shared.models.schemas.job_metadata import JobMetadataHelper
from shared.services.ai.llm_overrides import cleanup_llm_overrides, init_llm_overrides
from shared.services.ai.token_tracking import cleanup_token_tracker, init_token_tracker
from shared.services.jobs.lifecycle.service import get_sync_job_lifecycle_service
from shared.services.redis.distributed_lock import RedisJobLock
Expand Down Expand Up @@ -87,6 +88,7 @@ def _run_parse_job(
lifecycle_service.update_progress(job_id, progress=10, message="Parsing document...")
token_usage_dict = init_token_tracker()
stage_timing_dict = init_stage_tracker()
init_llm_overrides(JobMetadataHelper.get_llm_config(job_context.job_metadata))

try:
prepared_source = prepare_source_file(
Expand Down Expand Up @@ -180,6 +182,7 @@ def _run_parse_job(
"timing_ms": dict(stage_timing_dict),
"token_usage": dict(token_usage_dict),
}
cleanup_llm_overrides()
cleanup_token_tracker()
cleanup_stage_tracker()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from openai.types.chat import ChatCompletionMessageParam

from app.services.common.file_utils import path_handle
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client


def generate_fragment_title(content: str, max_tokens: int = 30) -> Optional[str]:
Expand All @@ -39,7 +39,8 @@ def generate_fragment_title(content: str, max_tokens: int = 30) -> Optional[str]
messages: list[ChatCompletionMessageParam] = [
{"role": "user", "content": title_prompt}
]
generated_title = get_openai_client().chat_completion(
client, _ = get_text_client()
generated_title = client.chat_completion(
messages=messages,
max_tokens=max_tokens,
timeout=30,
Expand Down
10 changes: 5 additions & 5 deletions apps/worker/app/services/document_parser/formats/image/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,8 @@
from app.services.common.file_loading import is_remote, load_file_bytes
from app.services.common.file_utils import path_handle
from shared.services.ai.summary.engine import summarize, transcribe
from shared.services.ai.openai_compatible_client_sync import (
OpenAICompatibleClientSync,
get_openai_client,
)
from shared.services.ai.llm_overrides import get_vision_client
from shared.services.ai.openai_compatible_client_sync import OpenAICompatibleClientSync

MD_IMAGE_PATTERN = r"!\[[^\]]*?\]\((.*?\.(?:png|jpe?g|gif))\)"
g_img_lock = threading.Lock()
Expand Down Expand Up @@ -61,7 +59,8 @@ def perceptual_hash(data: bytes) -> str:
def _get_vision_client() -> OpenAICompatibleClientSync:
"""Create OpenAI-compatible client for vision models, auto-routing by IMAGE_MODEL name."""
image_model = settings.IMAGE_MODEL or "qwen3.6-flash"
return get_openai_client(model=image_model)
client, _ = get_vision_client(requested_model=image_model)
return client


def image_bytes_to_base64(img_data: bytes, ext: str) -> str:
Expand Down Expand Up @@ -151,6 +150,7 @@ def ask_image(
image_model = settings.IMAGE_MODEL_MAX or "qwen3.6-flash"

if len(urls_) > 0:
client, image_model = get_vision_client(requested_model=image_model)
prompt, temperature, top_p, max_tokens = build_prompt(
task=task, texts=title_text, query=query, paras={"max_tokens": max_tokens}
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -179,7 +179,7 @@ def run_merge_pre_pass(

# ── 3. Call LLM directly (bypass df2md — texts is already formatted) ──
from shared.services.ai.prompt_service import build_prompt
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client
from shared.services.ai.response_process_service import eval_response

try:
Expand All @@ -194,7 +194,8 @@ def run_merge_pre_pass(
{"role": "user", "content": prompt},
]
with stage_timer("heading.merge_pre_pass_llm", group_count=len(groups), model_name=model_name):
answer = get_openai_client(model=model_name).chat_completion(
client, model_name = get_text_client(requested_model=model_name)
answer = client.chat_completion(
messages=messages,
model=model_name,
max_tokens=max_tokens,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@
from shared.services.ai.response_process_service import eval_response

# ARQ dependency is removed, use Celery instead
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client

# ==================== Helper Functions ====================

Expand Down Expand Up @@ -230,7 +230,8 @@ def _is_candidate_id(val):
candidate_count=n_candidates,
max_tokens=max_tokens,
):
answer = get_openai_client(model=model_name).chat_completion(
client, model_name = get_text_client(requested_model=model_name)
answer = client.chat_completion(
messages=messages,
model=model_name,
max_tokens=max_tokens,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@

from shared.services.ai.prompt_service import build_prompt
from shared.services.ai.response_process_service import eval_response
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client


# ==================== Markdown TOC Detection Functions ====================
Expand Down Expand Up @@ -230,7 +230,9 @@ def llm_judge_toc_range(
model_name=model_name,
total_candidates=total_candidates,
):
answer = get_openai_client(model=model_name).chat_completion(
client, resolved_model = get_text_client(requested_model=model_name)
model_name = resolved_model or model_name
answer = client.chat_completion(
messages=messages,
model=model_name,
max_tokens=max_tokens,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

from shared.services.ai.prompt_service import build_prompt
from shared.services.ai.response_process_service import eval_response
from shared.services.ai.openai_compatible_client_sync import get_openai_client
from shared.services.ai.llm_overrides import get_text_client
from shared.utils.text_utils import remove_duplicates_orderkept


Expand Down Expand Up @@ -68,7 +68,8 @@ def parse_headers_nonsmart(candidate_frame: pd.DataFrame) -> list[int]:
ttl=7200,
)

header_response = get_openai_client().chat_completion(
client, _ = get_text_client()
header_response = client.chat_completion(
messages=messages,
timeout=60,
usage_task="parser.table_detect_headers",
Expand Down
Loading
Loading