Files
tabby-web/backend/tabby/app/api.py

202 lines
5.6 KiB
Python
Raw Normal View History

2021-07-11 22:28:50 +02:00
import asyncio
2021-07-11 16:13:18 +02:00
import random
2021-06-15 23:43:23 +02:00
from django.conf import settings
from django.contrib.auth import logout
2021-07-11 22:28:50 +02:00
from dataclasses import dataclass
from pathlib import Path
2021-07-16 20:25:00 +02:00
from rest_framework import fields, status
from rest_framework.exceptions import APIException, PermissionDenied, NotFound
2021-06-15 23:43:23 +02:00
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
2021-06-25 21:55:40 +02:00
from rest_framework.mixins import ListModelMixin, RetrieveModelMixin, UpdateModelMixin
2021-06-15 23:43:23 +02:00
from rest_framework.views import APIView
from rest_framework.viewsets import GenericViewSet, ModelViewSet
2021-06-28 09:32:13 +02:00
from rest_framework.serializers import ModelSerializer, Serializer
2021-06-15 23:43:23 +02:00
from rest_framework_dataclasses.serializers import DataclassSerializer
2021-07-22 21:34:05 +02:00
from social_django.models import UserSocialAuth
2021-07-11 22:28:50 +02:00
from typing import List
2021-06-14 23:00:05 +02:00
2021-07-22 21:34:05 +02:00
from .consumers import GatewayAdminConnection
2021-07-24 15:48:12 +02:00
from .sponsors import check_is_sponsor, check_is_sponsor_cached
2021-07-11 16:13:18 +02:00
from .models import Config, Gateway, User
2021-06-14 23:00:05 +02:00
2021-06-15 23:43:23 +02:00
@dataclass
class AppVersion:
version: str
2021-07-11 22:28:50 +02:00
plugins: List[str]
2021-06-15 23:43:23 +02:00
class AppVersionSerializer(DataclassSerializer):
class Meta:
dataclass = AppVersion
2021-07-11 16:13:18 +02:00
class GatewaySerializer(ModelSerializer):
url = fields.SerializerMethodField()
2021-07-11 22:28:50 +02:00
auth_token = fields.CharField()
2021-07-11 16:13:18 +02:00
class Meta:
fields = '__all__'
model = Gateway
def get_url(self, gw):
return f'{"wss" if gw.secure else "ws"}://{gw.host}:{gw.port}/'
2021-06-14 23:00:05 +02:00
class ConfigSerializer(ModelSerializer):
2021-07-24 15:48:12 +02:00
name = fields.CharField(required=False)
2021-06-14 23:00:05 +02:00
class Meta:
model = Config
2021-06-15 23:43:23 +02:00
read_only_fields = ('user', 'created_at', 'modified_at')
2021-06-14 23:00:05 +02:00
fields = '__all__'
class ConfigViewSet(ModelViewSet):
queryset = Config.objects.all()
serializer_class = ConfigSerializer
2021-06-15 23:43:23 +02:00
permission_classes = [IsAuthenticated]
2021-06-14 23:00:05 +02:00
def get_queryset(self):
2021-06-15 23:43:23 +02:00
if self.request.user.is_authenticated:
return Config.objects.filter(user=self.request.user)
return Config.objects.none()
def perform_create(self, serializer):
serializer.save(user=self.request.user)
class AppVersionViewSet(ListModelMixin, GenericViewSet):
serializer_class = AppVersionSerializer
lookup_field = 'id'
lookup_value_regex = r'[\w\d.-]+'
queryset = ''
def _get_versions(self):
2021-07-11 22:28:50 +02:00
return [self._get_version(x) for x in settings.APP_DIST_PATH.iterdir()]
def _get_version(self, dir: Path):
plugins = [
x.name for x in dir.iterdir()
if x.is_dir() and x.name not in [
'tabby-web-container',
'tabby-web-demo',
]
]
return AppVersion(
version=dir.name,
plugins=plugins,
)
2021-06-15 23:43:23 +02:00
def list(self, request, *args, **kwargs):
2021-07-11 22:28:50 +02:00
return Response(
self.serializer_class(
self._get_versions(),
many=True,
).data
)
2021-06-15 23:43:23 +02:00
2021-06-16 23:49:44 +02:00
class UserSerializer(ModelSerializer):
id = fields.IntegerField()
2021-06-25 21:55:40 +02:00
is_pro = fields.SerializerMethodField()
2021-07-24 15:48:12 +02:00
is_sponsor = fields.SerializerMethodField()
2021-07-22 21:34:05 +02:00
github_username = fields.SerializerMethodField()
2021-06-16 23:49:44 +02:00
class Meta:
model = User
2021-07-11 22:28:50 +02:00
fields = (
'id',
'username',
'active_config',
'custom_connection_gateway',
'custom_connection_gateway_token',
2021-07-24 15:48:12 +02:00
'config_sync_token',
2021-07-11 22:28:50 +02:00
'is_pro',
2021-07-24 15:48:12 +02:00
'is_sponsor',
2021-07-22 21:34:05 +02:00
'github_username',
2021-07-11 22:28:50 +02:00
)
2021-06-16 23:49:44 +02:00
read_only_fields = ('id', 'username')
2021-06-25 21:55:40 +02:00
def get_is_pro(self, obj):
2021-07-24 15:48:12 +02:00
return check_is_sponsor_cached(obj) or obj.force_pro
def get_is_sponsor(self, obj):
return check_is_sponsor_cached(obj)
2021-07-22 21:34:05 +02:00
def get_github_username(self, obj):
social_auth = UserSocialAuth.objects.filter(user=obj, provider='github').first()
if not social_auth:
return None
return social_auth.extra_data.get('login')
2021-06-16 23:49:44 +02:00
2021-06-25 21:55:40 +02:00
class UserViewSet(RetrieveModelMixin, UpdateModelMixin, GenericViewSet):
2021-06-16 23:49:44 +02:00
queryset = User.objects.all()
serializer_class = UserSerializer
def get_object(self):
if self.request.user.is_authenticated:
return self.request.user
2021-06-25 21:55:40 +02:00
raise PermissionDenied()
2021-06-16 23:49:44 +02:00
2021-06-15 23:43:23 +02:00
class LogoutView(APIView):
def post(self, request, format=None):
logout(request)
2021-06-16 23:49:44 +02:00
return Response(None)
2021-06-28 09:32:13 +02:00
class InstanceInfoSerializer(Serializer):
login_enabled = fields.BooleanField()
class InstanceInfoViewSet(RetrieveModelMixin, GenericViewSet):
queryset = '' # type: ignore
serializer_class = InstanceInfoSerializer
def get_object(self):
return {
'login_enabled': settings.ENABLE_LOGIN,
}
2021-07-11 16:13:18 +02:00
2021-07-16 20:25:00 +02:00
class NoGatewaysError(APIException):
status_code = status.HTTP_503_SERVICE_UNAVAILABLE
default_detail ='No connection gateways available.'
default_code = 'no_gateways'
2021-07-11 16:13:18 +02:00
class ChooseGatewayViewSet(RetrieveModelMixin, GenericViewSet):
queryset = Gateway.objects.filter(enabled=True)
serializer_class = GatewaySerializer
2021-07-11 22:28:50 +02:00
async def _authorize_client(self, gw):
c = GatewayAdminConnection(gw)
await c.connect()
token = await c.authorize_client()
await c.close()
return token
2021-07-11 16:13:18 +02:00
def get_object(self):
gateways = list(self.queryset)
2021-07-16 20:25:00 +02:00
random.shuffle(gateways)
2021-07-11 16:13:18 +02:00
if not len(gateways):
raise NotFound()
2021-07-11 22:28:50 +02:00
loop = asyncio.new_event_loop()
2021-07-16 20:25:00 +02:00
try:
for gw in gateways:
try:
gw.auth_token = loop.run_until_complete(self._authorize_client(gw))
except ConnectionError:
continue
return gw
raise NoGatewaysError()
finally:
loop.close()