Merge branch 'rajat/dev-780' into vineeth/dev-769
This commit is contained in:
commit
9096e33b15
|
|
@ -107,8 +107,12 @@ def run_migrations_online() -> None:
|
|||
connection.execute(
|
||||
text(f"GRANT ALL ON SCHEMA {target_metadata.schema} TO current_user")
|
||||
)
|
||||
# Install pgvector extension if it doesn't exist
|
||||
connection.execute(text("CREATE EXTENSION IF NOT EXISTS vector"))
|
||||
# Set and verify search_path
|
||||
connection.execute(text(f"SET search_path TO {target_metadata.schema}, public"))
|
||||
connection.execute(
|
||||
text(f"SET search_path TO {target_metadata.schema}, public, extensions")
|
||||
)
|
||||
connection.commit()
|
||||
|
||||
context.configure(
|
||||
|
|
|
|||
86
src/agent.py
86
src/agent.py
|
|
@ -152,7 +152,6 @@ async def chat(
|
|||
user_id: str,
|
||||
session_id: str,
|
||||
queries: str | list[str],
|
||||
db: AsyncSession,
|
||||
stream: bool = False,
|
||||
) -> schemas.DialecticResponse | MessageStreamManager:
|
||||
"""
|
||||
|
|
@ -178,6 +177,7 @@ async def chat(
|
|||
|
||||
# Setup phase - create resources we'll need for all operations
|
||||
|
||||
<<<<<<< HEAD
|
||||
# 1. Create embedding store
|
||||
collection = await crud.get_or_create_user_protected_collection(db, app_id, user_id)
|
||||
|
||||
|
|
@ -205,42 +205,70 @@ async def chat(
|
|||
latest_message = latest_messages.scalar_one_or_none()
|
||||
latest_message_id = latest_message.public_id if latest_message else None
|
||||
logger.debug(f"Latest user message ID: {latest_message_id}")
|
||||
=======
|
||||
# 1. Fetch latest user message & chat history
|
||||
async with tracked_db("chat.load_history") as db_history:
|
||||
stmt = (
|
||||
select(models.Message)
|
||||
.where(models.Message.app_id == app_id)
|
||||
.where(models.Message.user_id == user_id)
|
||||
.where(models.Message.session_id == session_id)
|
||||
.where(models.Message.is_user)
|
||||
.order_by(models.Message.id.desc())
|
||||
.limit(1)
|
||||
)
|
||||
latest_messages = await db_history.execute(stmt)
|
||||
latest_message = latest_messages.scalar_one_or_none()
|
||||
latest_message_id = latest_message.public_id if latest_message else None
|
||||
logger.debug(f"Latest user message ID: {latest_message_id}")
|
||||
>>>>>>> rajat/dev-780
|
||||
|
||||
# Get chat history for the session
|
||||
chat_history, _, _ = await history.get_summarized_history(
|
||||
db, session_id, summary_type=history.SummaryType.SHORT
|
||||
)
|
||||
if not chat_history:
|
||||
logger.warning(f"No chat history found for session {session_id}")
|
||||
chat_history = f"someone asked this about the user's message: {final_query}"
|
||||
logger.debug(f"IDs: {app_id}, {user_id}, {session_id}")
|
||||
message_count = len(chat_history.split("\n"))
|
||||
logger.debug(f"Retrieved chat history: {message_count} messages")
|
||||
chat_history, _, _ = await history.get_summarized_history(
|
||||
db_history, session_id, summary_type=history.SummaryType.SHORT
|
||||
)
|
||||
if not chat_history:
|
||||
logger.warning(f"No chat history found for session {session_id}")
|
||||
chat_history = f"someone asked this about the user's message: {final_query}"
|
||||
logger.debug(f"IDs: {app_id}, {user_id}, {session_id}")
|
||||
message_count = len(chat_history.split("\n"))
|
||||
logger.debug(f"Retrieved chat history: {message_count} messages")
|
||||
|
||||
# Run both long-term and short-term context retrieval concurrently
|
||||
logger.debug("Starting parallel tasks for context retrieval")
|
||||
long_term_task = get_long_term_facts(final_query, embedding_store)
|
||||
short_term_task = run_tom_inference(chat_history, session_id)
|
||||
# Run short-term inference and long-term facts in parallel
|
||||
async def fetch_long_term():
|
||||
async with tracked_db("chat.get_collection_and_facts") as db_embed:
|
||||
collection = await crud.get_or_create_user_protected_collection(
|
||||
db_embed, app_id, user_id
|
||||
)
|
||||
embedding_store = CollectionEmbeddingStore(
|
||||
db=db_embed,
|
||||
app_id=app_id,
|
||||
user_id=user_id,
|
||||
collection_id=collection.public_id, # type: ignore
|
||||
)
|
||||
facts = await get_long_term_facts(final_query, embedding_store)
|
||||
return facts
|
||||
|
||||
long_term_task = asyncio.create_task(fetch_long_term())
|
||||
short_term_task = asyncio.create_task(run_tom_inference(chat_history, session_id))
|
||||
|
||||
# Wait for both tasks to complete
|
||||
facts, tom_inference = await asyncio.gather(long_term_task, short_term_task)
|
||||
logger.debug(f"Retrieved {len(facts)} facts from long-term memory")
|
||||
logger.debug(f"TOM inference completed with {len(tom_inference)} characters")
|
||||
|
||||
# Generate a fresh user representation
|
||||
logger.debug("Generating user representation")
|
||||
user_representation = await generate_user_representation(
|
||||
app_id=app_id,
|
||||
user_id=user_id,
|
||||
session_id=session_id,
|
||||
chat_history=chat_history,
|
||||
tom_inference=tom_inference,
|
||||
facts=facts,
|
||||
embedding_store=embedding_store,
|
||||
db=db,
|
||||
message_id=latest_message_id,
|
||||
with_inference=False,
|
||||
)
|
||||
async with tracked_db("chat.generate_user_representation") as db_rep:
|
||||
user_representation = await generate_user_representation(
|
||||
app_id=app_id,
|
||||
user_id=user_id,
|
||||
session_id=session_id,
|
||||
chat_history=chat_history,
|
||||
tom_inference=tom_inference,
|
||||
facts=facts,
|
||||
db=db_rep,
|
||||
message_id=latest_message_id,
|
||||
with_inference=False,
|
||||
)
|
||||
logger.debug(
|
||||
f"User representation generated: {len(user_representation)} characters"
|
||||
)
|
||||
|
|
@ -430,7 +458,6 @@ async def generate_user_representation(
|
|||
chat_history: str,
|
||||
tom_inference: str,
|
||||
facts: list[str],
|
||||
embedding_store: CollectionEmbeddingStore,
|
||||
db: AsyncSession,
|
||||
message_id: Optional[str] = None,
|
||||
with_inference: bool = False,
|
||||
|
|
@ -477,7 +504,6 @@ async def generate_user_representation(
|
|||
chat_history=chat_history,
|
||||
session_id=session_id,
|
||||
facts=facts,
|
||||
embedding_store=embedding_store,
|
||||
user_representation=latest_representation,
|
||||
tom_inference=tom_inference,
|
||||
)
|
||||
|
|
|
|||
11
src/db.py
11
src/db.py
|
|
@ -53,11 +53,12 @@ def init_db():
|
|||
echo=os.getenv("SQL_DEBUG", "false").lower() == "true",
|
||||
)
|
||||
|
||||
# Create schema if it doesn't exist
|
||||
if table_schema:
|
||||
with sync_engine.connect() as connection:
|
||||
connection.execute(text(f'CREATE SCHEMA IF NOT EXISTS "{table_schema}"'))
|
||||
connection.commit()
|
||||
with sync_engine.connect() as connection:
|
||||
# Create schema if it doesn't exist
|
||||
connection.execute(text(f'CREATE SCHEMA IF NOT EXISTS "{table_schema}"'))
|
||||
# Install pgvector extension if it doesn't exist
|
||||
connection.execute(text("CREATE EXTENSION IF NOT EXISTS vector"))
|
||||
connection.commit()
|
||||
|
||||
# Run Alembic migrations
|
||||
alembic_cfg = Config("alembic.ini")
|
||||
|
|
|
|||
|
|
@ -9,8 +9,6 @@ from sentry_sdk.ai.monitoring import ai_track
|
|||
from src.utils import parse_xml_content
|
||||
from src.utils.model_client import ModelClient, ModelProvider
|
||||
|
||||
from .embeddings import CollectionEmbeddingStore
|
||||
|
||||
# Configure logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -29,7 +27,6 @@ MAX_FACT_DISTANCE = 0.85
|
|||
async def get_user_representation_long_term(
|
||||
chat_history: str,
|
||||
session_id: str,
|
||||
embedding_store: CollectionEmbeddingStore,
|
||||
user_representation: str = "None",
|
||||
tom_inference: str = "None",
|
||||
facts: Optional[list[str]] = None,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import logging
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
import re
|
||||
from contextlib import asynccontextmanager
|
||||
|
|
@ -175,7 +176,7 @@ async def global_exception_handler(request: Request, exc: Exception):
|
|||
async def track_request(request: Request, call_next):
|
||||
# Create a request ID that includes endpoint information
|
||||
# Remove any IDs from the path - updated regex for NanoIDs (21 chars, A-Za-z0-9_-)
|
||||
endpoint = re.sub(r'/[A-Za-z0-9_-]{21}', '', request.url.path).replace("/", "_")
|
||||
endpoint = re.sub(r"/[A-Za-z0-9_-]{21}", "", request.url.path).replace("/", "_")
|
||||
request_id = f"{request.method}:{endpoint}:{str(uuid.uuid4())[:8]}"
|
||||
|
||||
# Store in request state and context var
|
||||
|
|
|
|||
|
|
@ -216,7 +216,6 @@ async def chat(
|
|||
options: schemas.DialecticOptions = Body(
|
||||
..., description="Dialectic Endpoint Parameters"
|
||||
),
|
||||
db=db,
|
||||
):
|
||||
|
||||
"""Chat with the Dialectic API"""
|
||||
|
|
@ -226,7 +225,6 @@ async def chat(
|
|||
user_id=user_id,
|
||||
session_id=session_id,
|
||||
queries=options.queries,
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
|
||||
|
|
@ -238,7 +236,6 @@ async def chat(
|
|||
session_id=session_id,
|
||||
queries=options.queries,
|
||||
stream=True,
|
||||
db=db,
|
||||
)
|
||||
if type(stream) is AsyncMessageStreamManager:
|
||||
async with stream as stream_manager:
|
||||
|
|
|
|||
Loading…
Reference in New Issue