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