diff --git a/apps/agent/serializers_admin.py b/apps/agent/serializers_admin.py new file mode 100644 index 0000000..008901c --- /dev/null +++ b/apps/agent/serializers_admin.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 diff --git a/apps/agent/urls.py b/apps/agent/urls.py new file mode 100644 index 0000000..0431225 --- /dev/null +++ b/apps/agent/urls.py @@ -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)), +] diff --git a/apps/agent/views_admin.py b/apps/agent/views_admin.py new file mode 100644 index 0000000..b772e51 --- /dev/null +++ b/apps/agent/views_admin.py @@ -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) diff --git a/config/urls.py b/config/urls.py index 13a4766..57b7289 100644 --- a/config/urls.py +++ b/config/urls.py @@ -81,6 +81,7 @@ api_patterns = [ path('videos/', include('apps.video.urls')), path('article/', include('apps.article.urls')), path('podcast/', include('apps.podcast.urls')), + path('agent/', include('apps.agent.urls')), path('bookmarks/', include('apps.bookmark.urls')), path('calendar/', include('apps.dobodbi_calendar.urls')), path('blog/', include('apps.blog.urls')),