Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 43 additions & 7 deletions scratchattach/__init__.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,33 @@
from .cloud.cloud import CustomCloud, ScratchCloud, TwCloud, get_cloud, get_scratch_cloud, get_tw_cloud
from .cloud.cloud import (
CustomCloud,
ScratchCloud,
TwCloud,
get_cloud,
get_scratch_cloud,
get_tw_cloud,
)
from .cloud._base import BaseCloud, AnyCloud

from .eventhandlers.cloud_server import init_cloud_server
from .eventhandlers._base import BaseEventHandler
from .eventhandlers.cloud_server import (
init_cloud_server,
init_ssl_cloud_server,
TwCloudSocket,
TwCloudServer,
TwSSLCloudServer,
)
from .eventhandlers._base import BaseEventHandler, BaseCloudServer
from .eventhandlers.filterbot import Filterbot, HardFilter, SoftFilter, SpamFilter
from .eventhandlers.cloud_storage import Database
from .eventhandlers.combine import MultiEventHandler

from .other.other_apis import *

# from .other.project_json_capabilities import ProjectBody, get_empty_project_pb, get_pb_from_dict, read_sb3_file, download_asset
# from .other.project_json_capabilities import (
# ProjectBody,
# get_empty_project_pb,
# get_pb_from_dict,
# read_sb3_file,
# download_asset,
# )
from .utils.encoder import Encoding
from .utils.enums import Languages, TTSVoices
from .utils.exceptions import (
Expand All @@ -27,11 +45,29 @@
from .site.cloud_activity import CloudActivity
from .site.forum import ForumPost, ForumTopic, get_topic, get_topic_list, youtube_link_to_scratch
from .site.project import Project, get_project, search_projects, explore_projects
from .site.session import Session, login, login_by_id, login_by_session_string, login_by_io, login_by_file, login_from_browser
from .site.session import (
Session,
login,
login_by_id,
login_by_session_string,
login_by_io,
login_by_file,
login_from_browser,
)
from .site.studio import Studio, get_studio, search_studios, explore_studios
from .site.classroom import Classroom, get_classroom
from .site.user import User, get_user, Rank
from .site._base import BaseSiteComponent
from .site.browser_cookies import Browser, ANY, FIREFOX, CHROME, CHROMIUM, VIVALDI, EDGE, EDGE_DEV, SAFARI
from .site.browser_cookies import (
Browser,
ANY,
FIREFOX,
CHROME,
CHROMIUM,
VIVALDI,
EDGE,
EDGE_DEV,
SAFARI,
)

from . import editor
207 changes: 203 additions & 4 deletions scratchattach/eventhandlers/_base.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,17 @@
from __future__ import annotations

import json
import time
import ssl
from abc import ABC, abstractmethod
from typing import Optional
from typing import Optional, Any
from collections import defaultdict
from threading import Thread, Event
from collections.abc import Callable
import traceback

from SimpleWebSocketServer import WebSocket

from scratchattach.utils.requests import requests
from scratchattach.utils import exceptions

Expand Down Expand Up @@ -41,7 +47,7 @@ def start(self, *, thread=True, ignore_exceptions=True):
else:
self._thread = None
self._updater()

def call_event(self, event_name, args : list = []):
try:
# print(f"Calling for {event_name}...")
Expand Down Expand Up @@ -69,7 +75,7 @@ def call_event(self, event_name, args : list = []):
@abstractmethod
def _updater(self):
pass

def __del__(self):
self.stop()

Expand Down Expand Up @@ -120,4 +126,197 @@ def inner(function):
return inner
else:
# => the decorator doesn't provide arguments
inner(function)
inner(function)

class BaseCloudServer(BaseEventHandler):
Comment thread
Boss-1s marked this conversation as resolved.
"""
Base class for all sa cloud servers.

If you are developing a custom cloud server with sa, please inherit from this class
and change up the methods as needed.
"""

hostname: str
"IP address or domain name of the host to bind the server to."
port: int
"Port to bind the server to."
tw_clients: dict[tuple[str, int], dict[str, Any]]
"Dictionary containing client information."
tw_variables: dict[str, dict[str, Any]]
"Dictionary containing existing cloud variables."
allow_non_numeric: bool
"Whether or not non-numeric characters are allowed in cloud variable values."
whitelisted_projects: list[str] | None
"Optional list of whitelisted projects."
length_limit: int | None
"Optional limit on the length of cloud variable values."
allow_nonscratch_names: bool
"Whether or not usernames that do not exist on scratch are allowed."
blocked_ips: list[str]
"List of blocked IP addresses."
sync_players: bool
log_var_sets: bool

def __init__(self,
hostname: str,
*,
port: int,
websocketclass: type[WebSocket],
length_limit: int | None = None,
allow_non_numeric: bool = True,
whitelisted_projects: list[Any] | None = None,
allow_nonscratch_names: bool = True,
blocked_ips: list[str] | None = None,
sync_players: bool = True,
log_var_sets: bool = True
):

if blocked_ips is None:
blocked_ips = []

BaseEventHandler.__init__(self)

self.running = False
self._events = {} # saves event functions called on cloud updates

self.tw_clients = {} # saves connected clients
self.tw_variables = {} # holds cloud variable states

self.hostname = hostname
self.port = port

# server config
self.allow_non_numeric = allow_non_numeric
self.whitelisted_projects = whitelisted_projects
self.length_limit = length_limit
self.allow_nonscratch_names = allow_nonscratch_names
self.blocked_ips = blocked_ips
self.sync_players = sync_players
self.log_var_sets = log_var_sets

def check_for_ip_ban(self, client):
if (
client.address[0] in self.blocked_ips
or client.address[0] + ":" + str(client.address[1]) in self.blocked_ips
or client.address in self.blocked_ips
):
client.sendMessage("You have been banned from this server")
client.close(4002)
print(client.address[0] + ":" + str(client.address[1]), "(IP-banned) was disconnected")
return True
return False

def active_projects(self):
only_active = {}
for project_id in self.tw_variables:
if self.active_user_ips(project_id) != []:
only_active[project_id] = self.tw_variables[project_id]
return only_active

def active_user_names(self, project_id):
return [self.tw_clients[user]["username"] for user in self.active_user_ips(project_id)]

def active_user_ips(self, project_id):
return list(filter(lambda user: str(self.tw_clients[user]["project_id"]) == str(project_id), self.tw_clients))

def get_global_vars(self):
return self.tw_variables

def get_project_vars(self, project_id):
project_id = str(project_id)
if project_id in self.tw_variables:
return self.tw_variables[project_id]
else:
return {}

def get_var(self, project_id, var_name):
project_id = str(project_id)
var_name = var_name.replace("☁ ", "")
if project_id in self.tw_variables:
if var_name in self.tw_variables[project_id]:
return self.tw_variables[project_id][var_name]
else:
return None
else:
return None

def set_global_vars(self, data):
for project_id in data:
self.set_project_vars(project_id, data[project_id])

def set_project_vars(self, project_id, data, *, user="@server"):
project_id = str(project_id)
self.tw_variables[project_id] = data
for client in (self.tw_clients[ip]["client"] for ip in self.active_user_ips(project_id)):
client.sendMessage(
"\n".join(
[
json.dumps(
{
"method": "set",
"project_id": project_id,
"name": "☁ " + varname,
"value": data[varname],
"server": "scratchattach/2.0.0",
"timestamp": time.time() * 1000,
"user": user,
}
)
for varname in data
]
)
)

def set_var(self, project_id, var_name, value, *, user="@server", skip_forward=None):
var_name = var_name.replace("☁ ", "")
project_id = str(project_id)
if project_id not in self.tw_variables:
self.tw_variables[project_id] = {}
self.tw_variables[project_id][var_name] = value

if self.sync_players is True:
for client in (self.tw_clients[ip]["client"] for ip in self.active_user_ips(project_id)):
if client == skip_forward:
continue
client.sendMessage(
json.dumps(
{
"method": "set",
"project_id": project_id,
"name": "☁ " + var_name,
"value": value,
"timestamp": time.time() * 1000,
"user": user,
}
)
)

def _check_value(self, value):
# Checks if a received cloud value satisfies the server's constraints
if self.length_limit is not None:
if len(str(value)) > self.length_limit:
return False
if self.allow_non_numeric is False:
x = value.replace(".", "")
x = x.replace("-", "")
if not (x.isnumeric() or x == ""):
return False
return True

def _updater(self):
try:
# Function called when .start() is executed (.start is inherited from BaseEventHandler)
print(f"Serving websocket server: ws://{self.hostname}:{self.port}")
self.serveforever()
except Exception as e:
raise exceptions.WebsocketServerError(str(e))

def pause(self):
self.running = False

def resume(self):
self.running = True

def stop(self, wait_call_threads: bool = True):
BaseEventHandler.stop(self, wait_call_threads)
self.close()
Loading