4 changed files with 136 additions and 0 deletions
-
58apps/agent/serializers_admin.py
-
17apps/agent/urls.py
-
60apps/agent/views_admin.py
-
1config/urls.py
@ -0,0 +1,58 @@ |
|||||
|
from rest_framework import serializers |
||||
|
|
||||
|
from .models import AgentPrompt, AgentSettings, EmbeddingSession |
||||
|
|
||||
|
|
||||
|
class AgentPromptSerializer(serializers.ModelSerializer): |
||||
|
id = serializers.IntegerField(required=False) |
||||
|
|
||||
|
class Meta: |
||||
|
model = AgentPrompt |
||||
|
fields = ["id", "content", "is_active"] |
||||
|
|
||||
|
|
||||
|
class AgentSettingsSerializer(serializers.ModelSerializer): |
||||
|
prompts = AgentPromptSerializer(many=True) |
||||
|
|
||||
|
class Meta: |
||||
|
model = AgentSettings |
||||
|
fields = ["id", "updated_at", "prompts"] |
||||
|
read_only_fields = ["id", "updated_at"] |
||||
|
|
||||
|
def update(self, instance, validated_data): |
||||
|
prompts_data = validated_data.pop("prompts", []) |
||||
|
instance = super().update(instance, validated_data) |
||||
|
|
||||
|
existing_prompts = {prompt.id: prompt for prompt in instance.prompts.all()} |
||||
|
kept_prompt_ids: list[int] = [] |
||||
|
|
||||
|
for prompt_data in prompts_data: |
||||
|
prompt_id = prompt_data.pop("id", None) |
||||
|
if prompt_id and prompt_id in existing_prompts: |
||||
|
prompt = existing_prompts[prompt_id] |
||||
|
prompt.content = prompt_data.get("content", prompt.content) |
||||
|
prompt.is_active = prompt_data.get("is_active", prompt.is_active) |
||||
|
prompt.save(update_fields=["content", "is_active"]) |
||||
|
kept_prompt_ids.append(prompt.id) |
||||
|
else: |
||||
|
prompt = AgentPrompt.objects.create(settings=instance, **prompt_data) |
||||
|
kept_prompt_ids.append(prompt.id) |
||||
|
|
||||
|
instance.prompts.exclude(id__in=kept_prompt_ids).delete() |
||||
|
instance.save() |
||||
|
return instance |
||||
|
|
||||
|
|
||||
|
class EmbeddingSessionSerializer(serializers.ModelSerializer): |
||||
|
class Meta: |
||||
|
model = EmbeddingSession |
||||
|
fields = [ |
||||
|
"id", |
||||
|
"status", |
||||
|
"progress", |
||||
|
"processed_items", |
||||
|
"total_items", |
||||
|
"error_message", |
||||
|
"created_at", |
||||
|
] |
||||
|
read_only_fields = fields |
||||
@ -0,0 +1,17 @@ |
|||||
|
from django.urls import include, path |
||||
|
from rest_framework.routers import SimpleRouter |
||||
|
|
||||
|
from .views_admin import ( |
||||
|
AdminAgentSettingsView, |
||||
|
AdminEmbeddingSessionCreateView, |
||||
|
AdminEmbeddingSessionViewSet, |
||||
|
) |
||||
|
|
||||
|
router = SimpleRouter() |
||||
|
router.register(r'embedding-sessions', AdminEmbeddingSessionViewSet, basename='admin-embedding-sessions') |
||||
|
|
||||
|
urlpatterns = [ |
||||
|
path('admin/settings/', AdminAgentSettingsView.as_view(), name='admin-agent-settings'), |
||||
|
path('admin/embedding-sessions/create/', AdminEmbeddingSessionCreateView.as_view(), name='admin-embedding-session-create'), |
||||
|
path('admin/', include(router.urls)), |
||||
|
] |
||||
@ -0,0 +1,60 @@ |
|||||
|
import threading |
||||
|
|
||||
|
import requests |
||||
|
from rest_framework import status |
||||
|
from rest_framework.authentication import TokenAuthentication |
||||
|
from rest_framework.generics import RetrieveUpdateAPIView |
||||
|
from rest_framework.permissions import IsAuthenticated |
||||
|
from rest_framework.response import Response |
||||
|
from rest_framework.views import APIView |
||||
|
from rest_framework.viewsets import ReadOnlyModelViewSet |
||||
|
|
||||
|
from apps.account.permissions import IsSuperAdmin |
||||
|
|
||||
|
from .models import AgentSettings, EmbeddingSession |
||||
|
from .serializers_admin import AgentSettingsSerializer, EmbeddingSessionSerializer |
||||
|
|
||||
|
|
||||
|
def trigger_embedding_session(session_id: int): |
||||
|
try: |
||||
|
requests.post( |
||||
|
"http://88.99.212.243:8098/api/sync-knowledge", |
||||
|
json={"session_id": session_id}, |
||||
|
timeout=5, |
||||
|
) |
||||
|
except Exception as exc: |
||||
|
session = EmbeddingSession.objects.filter(pk=session_id).first() |
||||
|
if session and session.status == "PENDING": |
||||
|
session.status = "FAILED" |
||||
|
session.error_message = str(exc) |
||||
|
session.save(update_fields=["status", "error_message"]) |
||||
|
|
||||
|
|
||||
|
class AdminAgentSettingsView(RetrieveUpdateAPIView): |
||||
|
serializer_class = AgentSettingsSerializer |
||||
|
permission_classes = [IsAuthenticated, IsSuperAdmin] |
||||
|
authentication_classes = [TokenAuthentication] |
||||
|
|
||||
|
def get_object(self): |
||||
|
return AgentSettings.load() |
||||
|
|
||||
|
|
||||
|
class AdminEmbeddingSessionViewSet(ReadOnlyModelViewSet): |
||||
|
serializer_class = EmbeddingSessionSerializer |
||||
|
permission_classes = [IsAuthenticated, IsSuperAdmin] |
||||
|
authentication_classes = [TokenAuthentication] |
||||
|
pagination_class = None |
||||
|
|
||||
|
def get_queryset(self): |
||||
|
return EmbeddingSession.objects.all().order_by("-created_at", "-id") |
||||
|
|
||||
|
|
||||
|
class AdminEmbeddingSessionCreateView(APIView): |
||||
|
permission_classes = [IsAuthenticated, IsSuperAdmin] |
||||
|
authentication_classes = [TokenAuthentication] |
||||
|
|
||||
|
def post(self, request): |
||||
|
session = EmbeddingSession.objects.create() |
||||
|
threading.Thread(target=trigger_embedding_session, args=(session.id,), daemon=True).start() |
||||
|
serializer = EmbeddingSessionSerializer(session, context={"request": request}) |
||||
|
return Response(serializer.data, status=status.HTTP_201_CREATED) |
||||
Write
Preview
Loading…
Cancel
Save
Reference in new issue