Commit dea78dc1 authored by Kantz's avatar Kantz
Browse files

alternatives retrival hinzugefügt

parent 325e6061
...@@ -68,6 +68,6 @@ Retrieval settings: ...@@ -68,6 +68,6 @@ Retrieval settings:
## Testing ## Testing
python -m test.hint_test --chat-id draft_session_mlgmxxzc_avmjfb python -m test.hint_test --chat-id draft_session_mlgmxxzc_avmjfb
python -m test.vector_store_test --query "Was ist eine Teilmenge?" --k 8 --expand python -m test.retrieval_store_test --query "Was ist eine Teilmenge?" --k 8 --expand
python -m test.math_intent_test --input "Integrate x^2" --input "Was ist 2+2?" python -m test.math_intent_test --input "Integrate x^2" --input "Was ist 2+2?"
python -m test.decision_test --chat-id draft_session_mlgmxxzc_avmjfb python -m test.decision_test --chat-id draft_session_mlgmxxzc_avmjfb
\ No newline at end of file
MATHPIX_APP_ID="" MATHPIX_APP_ID=""
MATHPIX_APP_KEY="" MATHPIX_APP_KEY=""
OPENAI_BASE_URL="" OPENAI_BASE_URL="https://chat-ai.academiccloud.de/v1/"
OPENAI_API_KEY="" OPENAI_API_KEY=""
OPENAI_CHAT_MODEL="" OPENAI_CHAT_MODEL="mistral-large-3-675b-instruct-2512"
OPENAI_CHAT_TEMPERATURE="" OPENAI_CHAT_TEMPERATURE="0.2"
OPENAI_EMBED_MODEL="" OPENAI_EMBED_MODEL="e5-mistral-7b-instruct"
EMBEDDING_TYPE="" # "openai-like" or "sentence-transformers" EMBEDDING_TYPE="sentence-transformer" # "openai-like" or "sentence-transformer"
EMBEDDING_DIM="512"
SENTENCE_TRANSFORMER_MODEL="jinaai/jina-embeddings-v4"
ORCHESTRATOR="tutor" # "tutor" or "qa"
RETRIEVAL_IMPL="child" # "child" or "subsection"
POSTGRES_URL="" POSTGRES_URL=""
OLLAMA_URL="" OLLAMA_URL="http://localhost:11434"
OLLAMA_MODEL= "" OLLAMA_MODEL= "ministral-3"
OLLAMA_TEMPERATURE="" OLLAMA_TEMPERATURE="0.2"
FRONTEND_URL="" FRONTEND_URL="http://localhost:5173"
\ No newline at end of file \ No newline at end of file
...@@ -18,6 +18,13 @@ class EmbeddingSettings(BaseModel): ...@@ -18,6 +18,13 @@ class EmbeddingSettings(BaseModel):
def get_orchestrator() -> str: def get_orchestrator() -> str:
return os.getenv("ORCHESTRATOR", "qa").lower() return os.getenv("ORCHESTRATOR", "qa").lower()
def get_retrieval_impl() -> str:
value = os.getenv("RETRIEVAL_IMPL", "child").strip().lower()
if value in {"child", "subsection"}:
return value
return "child"
def get_embedding_settings() -> EmbeddingSettings: def get_embedding_settings() -> EmbeddingSettings:
embedding_type = os.getenv("EMBEDDING_TYPE", "openai-like") embedding_type = os.getenv("EMBEDDING_TYPE", "openai-like")
if embedding_type == "sentence-transformer": if embedding_type == "sentence-transformer":
......
...@@ -8,8 +8,8 @@ from app.deterministic_services import ( ...@@ -8,8 +8,8 @@ from app.deterministic_services import (
Source, Source,
context_store, context_store,
referenz_decoder, referenz_decoder,
retrieval_store,
tool_logging, tool_logging,
vector_store,
) )
from app.deterministic_services.embeddings import EmbeddingFactory from app.deterministic_services.embeddings import EmbeddingFactory
...@@ -46,7 +46,7 @@ def extract_user_messages(messages: list[dict]) -> list[str]: ...@@ -46,7 +46,7 @@ def extract_user_messages(messages: list[dict]) -> list[str]:
def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]: def retrieve_context(query_text: str, pg_url: str | None = None) -> List[Source]:
embedder = EmbeddingFactory.create(config.get_embedding_settings()) embedder = EmbeddingFactory.create(config.get_embedding_settings())
url = pg_url or config.get_postgres_url() url = pg_url or config.get_postgres_url()
return vector_store.retrieve( return retrieval_store.retrieve(
pg_url=url, pg_url=url,
embedder=embedder, embedder=embedder,
query=query_text, query=query_text,
......
from __future__ import annotations
from typing import List
import app.config as config
from app.deterministic_services import vector_store, vector_store_subsection
from app.deterministic_services.vector_store import EmbeddingLike, Source
def _use_subsection_retrieval() -> bool:
return config.get_retrieval_impl() == "subsection"
def retrieve(
pg_url: str,
embedder: EmbeddingLike,
query: str,
k: int = 4,
section_index: int | None = None,
subsection_index: int | None = None,
source_type_filter: list[str] | None = None,
expand_links: bool = True,
neighbor_expand: int = 0,
) -> List[Source]:
if _use_subsection_retrieval():
return vector_store_subsection.retrieve(
pg_url=pg_url,
embedder=embedder,
query=query,
k=k,
section_index=section_index,
subsection_index=subsection_index,
source_type_filter=source_type_filter,
expand_links=expand_links,
neighbor_expand=neighbor_expand,
)
return vector_store.retrieve(
pg_url=pg_url,
embedder=embedder,
query=query,
k=k,
section_index=section_index,
subsection_index=subsection_index,
source_type_filter=source_type_filter,
expand_links=expand_links,
neighbor_expand=neighbor_expand,
)
from __future__ import annotations
from typing import Any, Dict, List, Optional
import psycopg
from pgvector import Vector
from pgvector.psycopg import register_vector
from psycopg.rows import dict_row
from app.deterministic_services.vector_store import (
EmbeddingLike,
Source,
_retrivla_to_sources,
_row_to_retrieved,
embed_query,
)
def retrieve(
pg_url: str,
embedder: EmbeddingLike,
query: str,
k: int = 4,
section_index: Optional[int] = None,
subsection_index: Optional[int] = None,
source_type_filter: Optional[List[str]] = None,
expand_links: bool = False,
neighbor_expand: int = 0,
) -> List[Source]:
# Parameters kept for drop-in compatibility with child-level retrieve.
_ = expand_links
_ = neighbor_expand
qvec = Vector(embed_query(embedder, query))
where = ["doc_type = ANY(%(sub_doc_types)s)"]
params: Dict[str, Any] = {
"qvec": qvec,
"k": k,
"sub_doc_types": ["subsection", "chapter"],
}
if section_index is not None:
where.append("section_index = %(section_index)s")
params["section_index"] = section_index
if subsection_index is not None:
where.append("subsection_index = %(subsection_index)s")
params["subsection_index"] = subsection_index
if source_type_filter:
where.append("source_type = ANY(%(source_type_filter)s)")
params["source_type_filter"] = source_type_filter
where_sql = " AND ".join(where)
sql = f"""
SELECT
uid, doc_type,
section_index, subsection_index, child_index,
section_title, subsection_title, title, source_type,
path, markdown,
1 - (embedding <=> %(qvec)s) AS score
FROM docs
WHERE {where_sql}
ORDER BY embedding <=> %(qvec)s
LIMIT %(k)s;
"""
with psycopg.connect(pg_url, row_factory=dict_row) as conn:
register_vector(conn)
with conn.cursor() as cur:
cur.execute(sql, params)
rows = cur.fetchall()
subsections = [_row_to_retrieved(row, source_type="subsection") for row in rows]
return _retrivla_to_sources({"subsections_direct": subsections})
...@@ -18,7 +18,7 @@ if str(ROOT_DIR) not in sys.path: ...@@ -18,7 +18,7 @@ if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR)) sys.path.insert(0, str(ROOT_DIR))
from app.deterministic_services.embeddings import BaseEmbeddings, EmbeddingFactory from app.deterministic_services.embeddings import BaseEmbeddings, EmbeddingFactory
from app.deterministic_services import vector_store from app.deterministic_services import retrieval_store, vector_store
def build_embedder() -> BaseEmbeddings: def build_embedder() -> BaseEmbeddings:
...@@ -51,7 +51,7 @@ def cli_query(args: argparse.Namespace) -> None: ...@@ -51,7 +51,7 @@ def cli_query(args: argparse.Namespace) -> None:
pg_url = args.pg or os.getenv("POSTGRES_URL") pg_url = args.pg or os.getenv("POSTGRES_URL")
if not pg_url: if not pg_url:
raise ValueError("Missing POSTGRES_URL") raise ValueError("Missing POSTGRES_URL")
result = vector_store.retrieve( result = retrieval_store.retrieve(
pg_url=pg_url, pg_url=pg_url,
embedder=embedder, embedder=embedder,
query=args.q, query=args.q,
......
import argparse import argparse
from app import config from app import config
from app.deterministic_services.embeddings import EmbeddingFactory, OpenAILikeEmbeddings from app.deterministic_services.embeddings import EmbeddingFactory
from app.deterministic_services import vector_store from app.deterministic_services import retrieval_store
def _normalize_sources(result: object) -> list[vector_store.Source]:
if isinstance(result, list):
return result
if isinstance(result, dict):
groups = {k: v for k, v in result.items() if isinstance(v, list)}
return vector_store._retrivla_to_sources(groups)
raise TypeError("Unexpected retrieval result type")
def main() -> None: def main() -> None:
...@@ -30,7 +21,7 @@ def main() -> None: ...@@ -30,7 +21,7 @@ def main() -> None:
pg_url = args.pg or config.get_postgres_url() pg_url = args.pg or config.get_postgres_url()
embedder = EmbeddingFactory.create(config.get_embedding_settings()) embedder = EmbeddingFactory.create(config.get_embedding_settings())
result = vector_store.retrieve( sources = retrieval_store.retrieve(
pg_url=pg_url, pg_url=pg_url,
embedder=embedder, embedder=embedder,
query=args.query, query=args.query,
...@@ -42,7 +33,6 @@ def main() -> None: ...@@ -42,7 +33,6 @@ def main() -> None:
neighbor_expand=args.neighbor_expand, neighbor_expand=args.neighbor_expand,
) )
sources = _normalize_sources(result)
if not sources: if not sources:
print("Keine Quellen gefunden.") print("Keine Quellen gefunden.")
return return
......
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment