"""
Conversation service - Business logic for conversation management.
Handles CRUD operations, context building, and summarization triggers.
"""
from uuid import UUID
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from src.domains.agents.base import get_agent
from src.domains.conversations.models import Conversation, Message
from src.domains.conversations.summarize import generate_summary
from src.shared.config import get_settings
from src.shared.logging import get_logger
from src.shared.tokens import count_tokens
logger = get_logger(__name__)
class ConversationService:
"""
Service for managing conversations and messages.
Handles:
- CRUD operations for conversations and messages
- Context building for agent prompts
- Automatic summarization when approaching token limits
"""
def __init__(self, session: AsyncSession):
"""
Initialize with database session.
Args:
session: Async SQLAlchemy session
"""
self.session = session
self.settings = get_settings()
# === Conversation CRUD ===
async def create(
self,
user_id: str,
agent_type: str = "explore",
working_dir: str = ".",
title: str | None = None,
) -> Conversation:
"""
Create a new conversation.
Args:
user_id: Owner's user ID
agent_type: Type of agent for this conversation
working_dir: Working directory for agent
title: Optional title (auto-generated from first message if None)
Returns:
Created Conversation object
"""
conversation = Conversation(
user_id=user_id,
agent_type=agent_type,
working_dir=working_dir,
title=title,
)
self.session.add(conversation)
await self.session.flush()
logger.info(f"Created conversation {conversation.id} for user {user_id}")
return conversation
async def get(self, conversation_id: UUID) -> Conversation | None:
"""Get conversation by ID without messages."""
result = await self.session.execute(
select(Conversation).where(Conversation.id == conversation_id)
)
return result.scalar_one_or_none()
async def get_with_messages(self, conversation_id: UUID) -> Conversation | None:
"""Get conversation by ID with messages loaded."""
result = await self.session.execute(
select(Conversation)
.options(selectinload(Conversation.messages))
.where(Conversation.id == conversation_id)
)
return result.scalar_one_or_none()
async def list_by_user(
self,
user_id: str,
limit: int = 50,
offset: int = 0,
) -> tuple[list[Conversation], int]:
"""
List conversations for a user.
Args:
user_id: User ID to filter by
limit: Maximum results to return
offset: Offset for pagination
Returns:
Tuple of (conversations, total_count)
"""
# Get total count
count_result = await self.session.execute(
select(func.count(Conversation.id))
.where(Conversation.user_id == user_id)
)
total = count_result.scalar() or 0
# Get conversations
result = await self.session.execute(
select(Conversation)
.where(Conversation.user_id == user_id)
.order_by(Conversation.updated_at.desc())
.limit(limit)
.offset(offset)
)
conversations = list(result.scalars().all())
return conversations, total
async def delete(self, conversation_id: UUID) -> bool:
"""Delete a conversation and all its messages."""
conversation = await self.get(conversation_id)
if conversation:
await self.session.delete(conversation)
logger.info(f"Deleted conversation {conversation_id}")
return True
return False
# === Message Operations ===
async def add_message(
self,
conversation_id: UUID,
role: str,
content: str,
) -> Message:
"""
Add a message to a conversation.
Args:
conversation_id: Conversation to add to
role: Message role (user, assistant, system, summary)
content: Message content
Returns:
Created Message object
"""
# Count tokens
token_count = count_tokens(content)
message = Message(
conversation_id=conversation_id,
role=role,
content=content,
token_count=token_count,
)
self.session.add(message)
# Update conversation total tokens
conversation = await self.get(conversation_id)
if conversation:
conversation.total_tokens += token_count
# Auto-generate title from first user message
if conversation.title is None and role == "user":
conversation.title = content[:100] + ("..." if len(content) > 100 else "")
await self.session.flush()
return message
# === Context Building ===
def build_context_prompt(
self,
messages: list[Message],
current_message: str,
) -> str:
"""
Build a prompt with conversation context.
Includes summary (if exists) and recent messages.
Args:
messages: All conversation messages
current_message: The current user message
Returns:
Formatted prompt with context
"""
parts = []
# Find most recent summary
summaries = [m for m in messages if m.is_summary]
if summaries:
latest_summary = summaries[-1]
parts.append(
f"\n{latest_summary.content}\n"
)
# Get recent non-summary messages
recent = [m for m in messages if not m.is_summary]
keep_count = self.settings.keep_recent_messages
recent = recent[-keep_count:] if len(recent) > keep_count else recent
if recent:
parts.append("")
for msg in recent:
role_label = msg.role.upper()
parts.append(f"{role_label}: {msg.content}")
parts.append("")
# Add current message
parts.append(f"\n{current_message}\n")
return "\n\n".join(parts)
# === Agent Integration ===
async def get_agent_response(
self,
conversation_id: UUID,
user_message: str,
) -> str:
"""
Get agent response with conversation context.
Args:
conversation_id: Conversation ID
user_message: Current user message
Returns:
Agent's response text
"""
conversation = await self.get_with_messages(conversation_id)
if not conversation:
raise ValueError(f"Conversation {conversation_id} not found")
agent = get_agent(conversation.agent_type)
if not agent:
raise ValueError(f"Unknown agent type: {conversation.agent_type}")
# Build context prompt
context_prompt = self.build_context_prompt(
conversation.messages,
user_message,
)
# Run agent
response = await agent.run(
context_prompt,
working_dir=conversation.working_dir,
)
return response
# === Summarization ===
async def should_summarize(self, conversation_id: UUID) -> bool:
"""
Check if conversation needs summarization.
Args:
conversation_id: Conversation to check
Returns:
True if summarization should be triggered
"""
conversation = await self.get(conversation_id)
if not conversation:
return False
threshold = self.settings.max_context_tokens * self.settings.summarization_threshold
return conversation.total_tokens > threshold
async def summarize_if_needed(self, conversation_id: UUID) -> bool:
"""
Summarize old messages if approaching token limit.
Args:
conversation_id: Conversation to check and potentially summarize
Returns:
True if summarization was performed
"""
if not await self.should_summarize(conversation_id):
return False
conversation = await self.get_with_messages(conversation_id)
if not conversation:
return False
messages = conversation.messages
keep_count = self.settings.keep_recent_messages
# Don't summarize if not enough messages
if len(messages) <= keep_count + 1:
return False
# Get messages to summarize (exclude recent and existing summaries)
non_summary_msgs = [m for m in messages if not m.is_summary]
to_summarize = non_summary_msgs[:-keep_count]
if not to_summarize:
return False
logger.info(
f"Summarizing {len(to_summarize)} messages in conversation {conversation_id}"
)
# Generate summary
summary_text = await generate_summary(
to_summarize,
working_dir=conversation.working_dir,
)
# Get ID of last summarized message
last_summarized_id = to_summarize[-1].id
# Calculate tokens being removed
removed_tokens = sum(m.token_count for m in to_summarize)
summary_tokens = count_tokens(summary_text)
# Add summary message
summary_message = Message(
conversation_id=conversation_id,
role="summary",
content=summary_text,
token_count=summary_tokens,
is_summary=True,
summarizes_up_to=last_summarized_id,
)
self.session.add(summary_message)
# Mark old messages as summarized (soft delete by excluding from context)
for msg in to_summarize:
msg.is_summary = True # Reuse flag to mark as "summarized away"
# Update conversation token count
conversation.total_tokens = conversation.total_tokens - removed_tokens + summary_tokens
await self.session.flush()
logger.info(
f"Summarization complete: removed {removed_tokens} tokens, "
f"added {summary_tokens} token summary"
)
return True