116 lines
4 KiB
Python
116 lines
4 KiB
Python
|
|
import uuid
|
|
from django.db import models
|
|
from rest_framework import viewsets, permissions, status
|
|
from rest_framework.decorators import action
|
|
from rest_framework.response import Response
|
|
|
|
from .models import AIInteraction, AgentConfig
|
|
from .serializers import (
|
|
AIInteractionSerializer, ChatRequestSerializer,
|
|
AgentConfigSerializer, AgentConfigWriteSerializer,
|
|
)
|
|
from .agent_service import AgentService
|
|
|
|
|
|
class AIInteractionViewSet(viewsets.ReadOnlyModelViewSet):
|
|
queryset = AIInteraction.objects.all()
|
|
serializer_class = AIInteractionSerializer
|
|
permission_classes = [permissions.IsAuthenticated]
|
|
|
|
def get_queryset(self):
|
|
qs = AIInteraction.objects.filter(user=self.request.user)
|
|
session = self.request.query_params.get("session_id")
|
|
if session:
|
|
qs = qs.filter(session_id=session)
|
|
return qs
|
|
|
|
@action(detail=False, methods=["post"])
|
|
def chat(self, request):
|
|
serializer = ChatRequestSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
|
|
agent_id = serializer.validated_data.get("agent_id")
|
|
message = serializer.validated_data["message"]
|
|
session_id = serializer.validated_data.get("session_id") or str(uuid.uuid4())
|
|
history = serializer.validated_data.get("history", [])
|
|
system_prompt_override = serializer.validated_data.get("system_prompt_override", "")
|
|
temperature = serializer.validated_data.get("temperature")
|
|
max_tokens = serializer.validated_data.get("max_tokens")
|
|
|
|
agent = None
|
|
if agent_id:
|
|
agent = AgentConfig.objects.filter(id=agent_id, is_active=True).first()
|
|
else:
|
|
agent = AgentConfig.objects.filter(is_default=True, is_active=True).first()
|
|
|
|
if not agent:
|
|
return Response(
|
|
{"error": "No active agent found. Create an agent or set one as default."},
|
|
status=status.HTTP_400_BAD_REQUEST,
|
|
)
|
|
|
|
AIInteraction.objects.create(
|
|
user=request.user,
|
|
session_id=session_id,
|
|
message_role=AIInteraction.MessageRole.USER,
|
|
message_content=message,
|
|
)
|
|
|
|
success, response_text, tokens_used = AgentService.execute_prompt(
|
|
config=agent,
|
|
prompt=message,
|
|
history=history,
|
|
system_prompt_override=system_prompt_override or None,
|
|
temperature=temperature,
|
|
max_tokens=max_tokens,
|
|
)
|
|
|
|
if not success:
|
|
return Response(
|
|
{"error": response_text, "session_id": session_id},
|
|
status=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
|
)
|
|
|
|
AIInteraction.objects.create(
|
|
user=request.user,
|
|
session_id=session_id,
|
|
message_role=AIInteraction.MessageRole.ASSISTANT,
|
|
message_content=response_text,
|
|
tokens_used=tokens_used or 0,
|
|
)
|
|
|
|
return Response({
|
|
"response": response_text,
|
|
"session_id": session_id,
|
|
"tokens_used": tokens_used,
|
|
"agent": agent.name,
|
|
})
|
|
|
|
@action(detail=False, methods=["get"])
|
|
def sessions(self, request):
|
|
sessions = (
|
|
AIInteraction.objects.filter(user=request.user)
|
|
.values("session_id")
|
|
.distinct()
|
|
.order_by("-created_at")[:20]
|
|
)
|
|
return Response([s["session_id"] for s in sessions if s["session_id"]])
|
|
|
|
|
|
class AgentConfigViewSet(viewsets.ModelViewSet):
|
|
queryset = AgentConfig.objects.all()
|
|
permission_classes = [permissions.IsAuthenticated]
|
|
|
|
def get_serializer_class(self):
|
|
if self.request.user.role in ("admin", "ops") or self.request.user.is_superuser:
|
|
return AgentConfigWriteSerializer
|
|
return AgentConfigSerializer
|
|
|
|
def get_queryset(self):
|
|
user = self.request.user
|
|
if user.role in ("admin", "ops") or user.is_superuser:
|
|
return AgentConfig.objects.all()
|
|
return AgentConfig.objects.filter(
|
|
models.Q(user=user) | models.Q(user__isnull=True, is_active=True)
|
|
)
|