git reimport

This commit is contained in:
2019-03-15 15:02:19 +04:00
commit 742797309a
90 changed files with 4411 additions and 0 deletions

40
src/main.py Normal file
View File

@@ -0,0 +1,40 @@
import asyncio
import logging
import yaml
from zvk.bot.bot import Bot
from zvk.util.paths import CONFIG_PATH
from zvk.util.zlogging import logger, formatter
def read_config():
try:
with open(CONFIG_PATH) as file:
return yaml.load(file)
except Exception:
logger.exception(f'Could not read config')
raise
def prod_logging():
info_file = logging.FileHandler('info.log')
info_file.setLevel(logging.INFO)
info_file.setFormatter(formatter)
logger.addHandler(info_file)
warning_file = logging.FileHandler('warning.log')
warning_file.setLevel(logging.WARNING)
warning_file.setFormatter(formatter)
logger.addHandler(warning_file)
if __name__ == '__main__':
prod_logging()
bot = Bot(config=read_config())
try:
asyncio.run(bot.run())
except KeyboardInterrupt:
logger.info('KeyboardInterrupt, shutting down')

0
src/zvk/__init__.py Normal file
View File

0
src/zvk/bot/__init__.py Normal file
View File

116
src/zvk/bot/bot.py Normal file
View File

@@ -0,0 +1,116 @@
import glob
from typing import Any, Dict
from zvk.bot.event_type import BotEventType
from zvk.bot.plugin import Plugin
from zvk.bot.trunk import Trunk
from zvk.event.event import Event
from zvk.event.queue import EventQueue
from zvk.plugins.vk.api import VKApi
from zvk.util.db import Database
from zvk.util.network import Network
from zvk.util.paths import PLUGIN_GLOB
from zvk.util.zlogging import logger
class Bot:
"""
An instance of a bot running on a specified account.
Attributes:
config: A config dict, holding information about keys, plugins, etc.
net: Network communication abstraction.
api: An interface to interact with VK servers.
event_queue: EventQueue that processes various events happening inside the bot.
"""
config: Dict[str, Any]
net: Network
api: VKApi
event_queue: EventQueue
plugins: Dict[str, Plugin]
db: Database
trunk: Trunk
def __init__(self, config):
"""
Initialize the bot, but do not start yet.
Args:
config: A config dict.
"""
self.config = config
self.net = Network(self.config)
self.api = VKApi(self.config, self.net)
self.db = Database(self.config['db_url'])
self.trunk = Trunk()
self.event_queue = EventQueue()
self.plugins = dict()
self._load_plugins()
def starting_env(self):
return dict(
bot=self,
config=self.config,
api=self.api,
net=self.net,
db=self.db,
trunk=self.trunk,
)
async def initialize(self) -> None:
self.net.initialize()
self.trunk.initialize()
async def run(self) -> bool:
"""
Asynchronously run the bot until the event queue finishes.
"""
logger.info('Initializing the bot')
await self.initialize()
logger.info('Initialization finished')
logger.info('Starting the bot')
await self.event_queue.run([Event(BotEventType.STARTUP, **self.starting_env())])
logger.info('Main queue finished, shutting down')
return not self.event_queue.is_dirty
def die(self) -> None:
"""
Schedules an end event to happen in the queue.
"""
logger.info('Suicide by forcibly stopping the queue')
self.event_queue.omae_wa_mou_shindeiru()
def _load_plugins(self):
paths = glob.glob(PLUGIN_GLOB, recursive=True)
whitelist = self.config['plugins']['whitelist']
blacklist = self.config['plugins']['blacklist']
for path in paths:
plugin = Plugin(path)
plugin.read()
if plugin.is_degenerate:
continue
self.plugins[plugin.name] = plugin
if plugin.name in blacklist:
logger.info(f'Plugin {plugin.name} is blacklisted')
continue
if whitelist and plugin.name not in whitelist:
logger.info(f'Plugin {plugin.name} is not whitelisted')
continue
plugin.activate(self.event_queue)
logger.info(f'{len(self.plugins)} plugins loaded')

View File

@@ -0,0 +1,5 @@
from enum import Enum, auto
class BotEventType(Enum):
STARTUP = auto()

64
src/zvk/bot/plugin.py Normal file
View File

@@ -0,0 +1,64 @@
import importlib
import re
from typing import Set
from zvk.event.consumer import EventConsumer
from zvk.event.queue import EventQueue
from zvk.util.zlogging import logger
class Plugin:
path: str
import_path: str
name: str
consumers: Set[EventConsumer]
is_degenerate: bool
is_activated: bool
def __init__(self, path):
self.path = path
self.import_path = self.path
self.import_path = re.sub(r'/', '.', self.import_path)
self.import_path = re.sub(r'.py$', '', self.import_path)
self.import_path = re.sub(r'^src\.', '', self.import_path)
self.name = re.sub(r'^zvk\.plugins\.', '', self.import_path)
self.consumers = set()
self.is_degenerate = False
self.is_activated = False
def read(self):
logger.info(f'Reading plugin {self.name}')
module = importlib.import_module(self.import_path)
for sub_name in dir(module):
sub = getattr(module, sub_name)
if isinstance(sub, EventConsumer):
self.consumers.add(sub)
logger.debug(f'Found a consumer {self.name} -> {sub_name} = {sub}')
if len(self.consumers) == 0:
self.is_degenerate = True
def activate(self, queue: EventQueue):
for consumer in self.consumers:
queue.register_consumer(consumer)
logger.info(f'Plugin {self.name} activated')
self.is_activated = True
def deactivate(self, queue: EventQueue):
for consumer in self.consumers:
queue.deregister_consumer(consumer)
logger.info(f'Plugin {self.name} deactivated')
self.is_activated = False
def __str__(self):
return self.name

26
src/zvk/bot/trunk.py Normal file
View File

@@ -0,0 +1,26 @@
import asyncio
from typing import Dict, Any
class Trunk:
contents: Dict[str, asyncio.Future]
loop: asyncio.AbstractEventLoop
def __init__(self):
self.contents = dict()
self.loop = None
def initialize(self) -> None:
self.loop = asyncio.get_running_loop()
def set(self, key, value) -> None:
if key not in self.contents:
self.contents[key] = self.loop.create_future()
self.contents[key].set_result(value)
async def get(self, key) -> Any:
if key not in self.contents:
self.contents[key] = self.loop.create_future()
return await self.contents[key]

View File

87
src/zvk/event/consumer.py Normal file
View File

@@ -0,0 +1,87 @@
from typing import List, AsyncGenerator
from zvk.bot.event_type import BotEventType
from zvk.event.event import EventType, CoroutineFactory, Event
from zvk.event.reflection import run_with_env
async def async_generator_adapter(coro) -> AsyncGenerator[Event, None]:
if hasattr(coro, '__anext__') and hasattr(coro, '__aiter__'):
# our coro is an async generator, let's forward the output events downstream
async for output_event in coro:
yield output_event
elif hasattr(coro, '__await__'):
# coro is a simple man, just run it
await coro
else:
raise ValueError(f'{coro} is not an awaitable at all')
class EventConsumer:
"""
Wrapper around a coroutine factory (async def ...) that feeds it events of declared types.
Attributes:
consumes: List of `EventType`s that this consumer is going to trigger on.
coroutine_factory: `Callable` that is going to be called when corresponding events happen.
"""
consumes: List[EventType]
coroutine_factory: CoroutineFactory
def __init__(self, consumes: List[EventType]):
self.consumes = consumes
self.coroutine_factory = None
def __call__(self, *args, **kwargs):
if self.coroutine_factory is None:
self.coroutine_factory = args[0]
return self
return self.coroutine_factory(*args, **kwargs)
async def consume(self, event: Event) -> AsyncGenerator[Event, None]:
"""
This is an async generator that consumes an event and optionally generates a sequence of other events.
Args:
event: Event to consume.
Yields:
Events produced by the consumer.
"""
coro = run_with_env(event.env, self.coroutine_factory)
async for event in async_generator_adapter(coro):
yield event
# def async_gen_adapter()
def event_consumer(consumes: List[str]) -> EventConsumer:
"""
Decorator for event consumers.
Args:
consumes: List of `EventType`s that this consumer wants.
Returns:
An initialized EventConsumer instance.
"""
if callable(consumes):
# direct decoration
raise TypeError('Direct decoration is forbidden')
# create the decorating object first
return EventConsumer(consumes=consumes)
def on_startup(func=None) -> EventConsumer:
consumer = EventConsumer(consumes=[BotEventType.STARTUP])
if callable(func):
return consumer(func)
return consumer

50
src/zvk/event/event.py Normal file
View File

@@ -0,0 +1,50 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Callable, Hashable, Any
EventType = Hashable
CoroutineFactory = Callable
@dataclass
class Event:
"""
Describes a pretty general event.
Attributes:
event_type: A reference for event consumers. Can be anything hashable.
env: Environment that the event consumers are going to run inside. Will be used for fancy parameter substitution.
"""
event_type: EventType
env: Dict[str, Any]
def __init__(self, event_type, **kwargs):
if not hasattr(event_type, '__hash__'):
raise ValueError(f'Bad event type {event_type}')
self.event_type = event_type
self.env = dict()
self.env['env'] = self.env
self.env['event'] = self
self.env['event_type'] = event_type
self.env.update(kwargs)
def prepopulate_env_from(self, previous_event: Event) -> None:
"""
If an event was produced by a consumer from another event, populate its environment with old values.
Args:
previous_event: Event that produced `self`.
"""
new_env = dict(previous_event.env)
new_env.update(self.env)
self.env = new_env
def __str__(self):
return f'Event(event_type={self.event_type}, env=#{len(self.env)})'

49
src/zvk/event/periodic.py Normal file
View File

@@ -0,0 +1,49 @@
import asyncio
from datetime import datetime, timedelta
from enum import Enum, auto
from zvk.bot.event_type import BotEventType
from zvk.event.consumer import EventConsumer
from zvk.event.event import Event
from zvk.util.zlogging import logger
class PeriodicEventConsumer(EventConsumer):
def __init__(self, period_secs):
class SpecificPeriodicTick(Enum):
TICK = auto()
self.tick_event_type = SpecificPeriodicTick.TICK
super().__init__(consumes=[
BotEventType.STARTUP,
SpecificPeriodicTick.TICK
])
self.period = timedelta(seconds=period_secs)
self.tick_started_at = None
async def consume(self, event: Event):
self.tick_started_at = datetime.utcnow()
async for output_event in super().consume(event):
yield output_event
exec_duration = datetime.utcnow() - self.tick_started_at
wait_duration = self.period - exec_duration
logger.debug(f'A periodic hook is going to sleep for {wait_duration}')
if wait_duration > timedelta():
await asyncio.sleep(wait_duration.total_seconds())
logger.debug(f'A periodic hook woke up')
yield Event(self.tick_event_type)
def periodic(period_secs) -> EventConsumer:
if callable(period_secs):
# direct decoration
raise TypeError('Direct decoration is forbidden')
return PeriodicEventConsumer(period_secs)

152
src/zvk/event/queue.py Normal file
View File

@@ -0,0 +1,152 @@
from __future__ import annotations
import asyncio
from typing import Set, Dict, List
from zvk.event.consumer import EventConsumer
from zvk.event.event import Event, EventType
from zvk.util.zlogging import logger
class Task:
queue: EventQueue
consumer: EventConsumer
event: Event
asyncio_task: asyncio.Task
is_finalized: bool
def __init__(self, queue, consumer, event):
self.queue = queue
self.consumer = consumer
self.event = event
self.asyncio_task = None
self.is_finalized = False
def finalize(self):
if self.is_finalized:
return
self.queue.all_running_tasks.remove(self)
self.queue.consumer_to_running_tasks[self.consumer].remove(self)
if len(self.queue.all_running_tasks) == 0:
self.queue.has_finished.set()
self.is_finalized = True
def schedule(self):
self.asyncio_task = asyncio.create_task(self.run())
self.queue.all_running_tasks.add(self)
self.queue.consumer_to_running_tasks[self.consumer].add(self)
self.queue.has_finished.clear()
async def run(self):
try:
output_events_generator = self.consumer.consume(self.event)
async for output_event in output_events_generator:
logger.debug(f'Process {self.event} -> {output_event}')
output_event.prepopulate_env_from(self.event)
self.queue._route_event(output_event)
except asyncio.CancelledError:
logger.warning(f'A meal was cancelled, perhaps we are shutting down?')
except Exception:
# TODO: consumer banning/revival
logger.exception(f'A consumer died eating his cake. What am I gonna do?')
logger.info(f'Killing offending consumer {self.consumer}')
self.queue.deregister_consumer(consumer=self.consumer)
self.queue.is_dirty = True
finally:
self.finalize()
def cancel(self):
self.asyncio_task.cancel()
self.finalize()
class EventQueue:
all_consumers: Set[EventConsumer]
event_type_to_consumers: Dict[EventType, Set[EventConsumer]]
all_running_tasks: Set[asyncio.Task]
consumer_to_running_tasks: Dict[EventConsumer, Set[asyncio.Task]]
is_dead: bool
has_finished: asyncio.Event
is_dirty: bool
def __init__(self):
self.all_consumers = set()
self.event_type_to_consumers = dict()
self.all_running_tasks = set()
self.consumer_to_running_tasks = dict()
self.is_dead = False
self.has_finished = None
self.is_dirty = False
def _route_event(self, event: Event):
logger.debug(f'Routing event {event}')
if event.event_type not in self.event_type_to_consumers:
logger.info(f'No consumers defined for event {event}')
return
for consumer in self.event_type_to_consumers[event.event_type]:
Task(self, consumer, event).schedule()
def register_consumer(self, consumer: EventConsumer) -> None:
if consumer in self.all_consumers:
raise ValueError(f'Consumer {consumer} is already registered.')
self.all_consumers.add(consumer)
# register tasks
self.consumer_to_running_tasks[consumer] = set()
# register hooks
for event_type in consumer.consumes:
self.event_type_to_consumers.setdefault(event_type, set()).add(consumer)
def deregister_consumer(self, consumer: EventConsumer) -> None:
if consumer not in self.all_consumers:
raise ValueError(f'Consumer {consumer} is not registered.')
self.all_consumers.remove(consumer)
# deregister tasks
for running_task in set(self.consumer_to_running_tasks[consumer]):
running_task.cancel()
del self.consumer_to_running_tasks[consumer]
# deregister hooks
for event_type in consumer.consumes:
self.event_type_to_consumers[event_type].remove(consumer)
if len(self.event_type_to_consumers[event_type]) == 0:
del self.event_type_to_consumers[event_type]
async def run(self, starting_events: List[Event]) -> None:
self.has_finished = asyncio.Event()
self.has_finished.set()
for event in starting_events:
event.env['event_queue'] = self
self._route_event(event)
await self.has_finished.wait()
def omae_wa_mou_shindeiru(self) -> None:
logger.warning(f'Shutting down the queue')
for consumer in set(self.all_consumers):
self.deregister_consumer(consumer)
assert len(self.all_consumers) == 0
assert len(self.all_running_tasks) == 0

View File

@@ -0,0 +1,24 @@
import inspect
from typing import Callable
def run_with_env(env: dict, f: Callable):
"""
Magic to run a function with arguments taken from a dict `env`.
:param env: environment with possible arguments.
:param f: function to run.
:return: execution result.
"""
signature = inspect.signature(f)
bind = {}
for name, parameter in signature.parameters.items():
if name in env:
bind[name] = env[name]
elif parameter.default is not inspect.Parameter.empty:
bind[name] = parameter.default
else:
raise TypeError(f'Cannot find desired parameter: {f} wants {name} from {env}')
return f(**bind)

0
src/zvk/misc/__init__.py Normal file
View File

View File

@@ -0,0 +1,427 @@
# Generated by the protocol buffer compiler. DO NOT EDIT!
# source: Timetable.proto
import sys
_b=sys.version_info[0]<3 and (lambda x:x) or (lambda x:x.encode('latin1'))
from google.protobuf import descriptor as _descriptor
from google.protobuf import message as _message
from google.protobuf import reflection as _reflection
from google.protobuf import symbol_database as _symbol_database
from google.protobuf import descriptor_pb2
# @@protoc_insertion_point(imports)
_sym_db = _symbol_database.Default()
DESCRIPTOR = _descriptor.FileDescriptor(
name='Timetable.proto',
package='',
serialized_pb=_b('\n\x0fTimetable.proto\"\xe2\x01\n\tTimetable\x12\x1a\n\nproperties\x18\x01 \x02(\x0b\x32\x06.Props\x12\x18\n\x07subject\x18\x02 \x03(\x0b\x32\x07.Record\x12\x18\n\x07teacher\x18\x03 \x03(\x0b\x32\x07.Record\x12\x16\n\x05place\x18\x04 \x03(\x0b\x32\x07.Record\x12\x15\n\x04kind\x18\x05 \x03(\x0b\x32\x07.Record\x12\x16\n\x05group\x18\x06 \x03(\x0b\x32\x07.Record\x12\x10\n\x08subgroup\x18\x07 \x03(\t\x12\x17\n\x06lesson\x18\x08 \x03(\x0b\x32\x07.Lesson\x12\x13\n\x04task\x18\t \x03(\x0b\x32\x05.Task\"T\n\x05Props\x12\x12\n\nterm_start\x18\x01 \x02(\x06\x12\x13\n\x0bterm_length\x18\x02 \x02(\x05\x12\x13\n\x0bweeks_count\x18\x03 \x02(\x05\x12\r\n\x05times\x18\x04 \x02(\t\"V\n\x06Record\x12\x0b\n\x03gid\x18\x01 \x02(\x06\x12\x0c\n\x04name\x18\x02 \x02(\t\x12\x11\n\tfull_name\x18\x03 \x01(\t\x12\x10\n\x08\x63olor_id\x18\x04 \x01(\x05\x12\x0c\n\x04link\x18\x05 \x01(\t\"\xc0\x01\n\x06Lesson\x12\x13\n\x0bsubgroup_id\x18\x01 \x01(\x05\x12\x0b\n\x03\x64\x61y\x18\x02 \x02(\x05\x12\x0c\n\x04time\x18\x03 \x02(\t\x12\r\n\x05weeks\x18\x04 \x02(\t\x12\x12\n\nsubject_id\x18\x05 \x02(\x05\x12\x0f\n\x07kind_id\x18\x06 \x01(\x05\x12\x10\n\x08place_id\x18\x07 \x01(\x05\x12\x16\n\nteacher_id\x18\x08 \x03(\x05\x42\x02\x10\x01\x12\x14\n\x08group_id\x18\t \x03(\x05\x42\x02\x10\x01\x12\x12\n\nno_silence\x18\x64 \x01(\x08\"b\n\x04Task\x12\x12\n\nsubject_id\x18\x01 \x02(\x05\x12\x11\n\tday_index\x18\x02 \x02(\x05\x12\r\n\x05title\x18\x03 \x02(\t\x12\x13\n\x0b\x64\x65scription\x18\x04 \x01(\t\x12\x0f\n\x07\x64one_at\x18\x05 \x01(\x06')
)
_sym_db.RegisterFileDescriptor(DESCRIPTOR)
_TIMETABLE = _descriptor.Descriptor(
name='Timetable',
full_name='Timetable',
filename=None,
file=DESCRIPTOR,
containing_type=None,
fields=[
_descriptor.FieldDescriptor(
name='properties', full_name='Timetable.properties', index=0,
number=1, type=11, cpp_type=10, label=2,
has_default_value=False, default_value=None,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='subject', full_name='Timetable.subject', index=1,
number=2, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='teacher', full_name='Timetable.teacher', index=2,
number=3, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='place', full_name='Timetable.place', index=3,
number=4, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='kind', full_name='Timetable.kind', index=4,
number=5, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='group', full_name='Timetable.group', index=5,
number=6, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='subgroup', full_name='Timetable.subgroup', index=6,
number=7, type=9, cpp_type=9, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='lesson', full_name='Timetable.lesson', index=7,
number=8, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='task', full_name='Timetable.task', index=8,
number=9, type=11, cpp_type=10, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
],
extensions=[
],
nested_types=[],
enum_types=[
],
options=None,
is_extendable=False,
extension_ranges=[],
oneofs=[
],
serialized_start=20,
serialized_end=246,
)
_PROPS = _descriptor.Descriptor(
name='Props',
full_name='Props',
filename=None,
file=DESCRIPTOR,
containing_type=None,
fields=[
_descriptor.FieldDescriptor(
name='term_start', full_name='Props.term_start', index=0,
number=1, type=6, cpp_type=4, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='term_length', full_name='Props.term_length', index=1,
number=2, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='weeks_count', full_name='Props.weeks_count', index=2,
number=3, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='times', full_name='Props.times', index=3,
number=4, type=9, cpp_type=9, label=2,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
],
extensions=[
],
nested_types=[],
enum_types=[
],
options=None,
is_extendable=False,
extension_ranges=[],
oneofs=[
],
serialized_start=248,
serialized_end=332,
)
_RECORD = _descriptor.Descriptor(
name='Record',
full_name='Record',
filename=None,
file=DESCRIPTOR,
containing_type=None,
fields=[
_descriptor.FieldDescriptor(
name='gid', full_name='Record.gid', index=0,
number=1, type=6, cpp_type=4, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='name', full_name='Record.name', index=1,
number=2, type=9, cpp_type=9, label=2,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='full_name', full_name='Record.full_name', index=2,
number=3, type=9, cpp_type=9, label=1,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='color_id', full_name='Record.color_id', index=3,
number=4, type=5, cpp_type=1, label=1,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='link', full_name='Record.link', index=4,
number=5, type=9, cpp_type=9, label=1,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
],
extensions=[
],
nested_types=[],
enum_types=[
],
options=None,
is_extendable=False,
extension_ranges=[],
oneofs=[
],
serialized_start=334,
serialized_end=420,
)
_LESSON = _descriptor.Descriptor(
name='Lesson',
full_name='Lesson',
filename=None,
file=DESCRIPTOR,
containing_type=None,
fields=[
_descriptor.FieldDescriptor(
name='subgroup_id', full_name='Lesson.subgroup_id', index=0,
number=1, type=5, cpp_type=1, label=1,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='day', full_name='Lesson.day', index=1,
number=2, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='time', full_name='Lesson.time', index=2,
number=3, type=9, cpp_type=9, label=2,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='weeks', full_name='Lesson.weeks', index=3,
number=4, type=9, cpp_type=9, label=2,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='subject_id', full_name='Lesson.subject_id', index=4,
number=5, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='kind_id', full_name='Lesson.kind_id', index=5,
number=6, type=5, cpp_type=1, label=1,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='place_id', full_name='Lesson.place_id', index=6,
number=7, type=5, cpp_type=1, label=1,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='teacher_id', full_name='Lesson.teacher_id', index=7,
number=8, type=5, cpp_type=1, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=_descriptor._ParseOptions(descriptor_pb2.FieldOptions(), _b('\020\001'))),
_descriptor.FieldDescriptor(
name='group_id', full_name='Lesson.group_id', index=8,
number=9, type=5, cpp_type=1, label=3,
has_default_value=False, default_value=[],
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=_descriptor._ParseOptions(descriptor_pb2.FieldOptions(), _b('\020\001'))),
_descriptor.FieldDescriptor(
name='no_silence', full_name='Lesson.no_silence', index=9,
number=100, type=8, cpp_type=7, label=1,
has_default_value=False, default_value=False,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
],
extensions=[
],
nested_types=[],
enum_types=[
],
options=None,
is_extendable=False,
extension_ranges=[],
oneofs=[
],
serialized_start=423,
serialized_end=615,
)
_TASK = _descriptor.Descriptor(
name='Task',
full_name='Task',
filename=None,
file=DESCRIPTOR,
containing_type=None,
fields=[
_descriptor.FieldDescriptor(
name='subject_id', full_name='Task.subject_id', index=0,
number=1, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='day_index', full_name='Task.day_index', index=1,
number=2, type=5, cpp_type=1, label=2,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='title', full_name='Task.title', index=2,
number=3, type=9, cpp_type=9, label=2,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='description', full_name='Task.description', index=3,
number=4, type=9, cpp_type=9, label=1,
has_default_value=False, default_value=_b("").decode('utf-8'),
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
_descriptor.FieldDescriptor(
name='done_at', full_name='Task.done_at', index=4,
number=5, type=6, cpp_type=4, label=1,
has_default_value=False, default_value=0,
message_type=None, enum_type=None, containing_type=None,
is_extension=False, extension_scope=None,
options=None),
],
extensions=[
],
nested_types=[],
enum_types=[
],
options=None,
is_extendable=False,
extension_ranges=[],
oneofs=[
],
serialized_start=617,
serialized_end=715,
)
_TIMETABLE.fields_by_name['properties'].message_type = _PROPS
_TIMETABLE.fields_by_name['subject'].message_type = _RECORD
_TIMETABLE.fields_by_name['teacher'].message_type = _RECORD
_TIMETABLE.fields_by_name['place'].message_type = _RECORD
_TIMETABLE.fields_by_name['kind'].message_type = _RECORD
_TIMETABLE.fields_by_name['group'].message_type = _RECORD
_TIMETABLE.fields_by_name['lesson'].message_type = _LESSON
_TIMETABLE.fields_by_name['task'].message_type = _TASK
DESCRIPTOR.message_types_by_name['Timetable'] = _TIMETABLE
DESCRIPTOR.message_types_by_name['Props'] = _PROPS
DESCRIPTOR.message_types_by_name['Record'] = _RECORD
DESCRIPTOR.message_types_by_name['Lesson'] = _LESSON
DESCRIPTOR.message_types_by_name['Task'] = _TASK
Timetable = _reflection.GeneratedProtocolMessageType('Timetable', (_message.Message,), dict(
DESCRIPTOR = _TIMETABLE,
__module__ = 'Timetable_pb2'
# @@protoc_insertion_point(class_scope:Timetable)
))
_sym_db.RegisterMessage(Timetable)
Props = _reflection.GeneratedProtocolMessageType('Props', (_message.Message,), dict(
DESCRIPTOR = _PROPS,
__module__ = 'Timetable_pb2'
# @@protoc_insertion_point(class_scope:Props)
))
_sym_db.RegisterMessage(Props)
Record = _reflection.GeneratedProtocolMessageType('Record', (_message.Message,), dict(
DESCRIPTOR = _RECORD,
__module__ = 'Timetable_pb2'
# @@protoc_insertion_point(class_scope:Record)
))
_sym_db.RegisterMessage(Record)
Lesson = _reflection.GeneratedProtocolMessageType('Lesson', (_message.Message,), dict(
DESCRIPTOR = _LESSON,
__module__ = 'Timetable_pb2'
# @@protoc_insertion_point(class_scope:Lesson)
))
_sym_db.RegisterMessage(Lesson)
Task = _reflection.GeneratedProtocolMessageType('Task', (_message.Message,), dict(
DESCRIPTOR = _TASK,
__module__ = 'Timetable_pb2'
# @@protoc_insertion_point(class_scope:Task)
))
_sym_db.RegisterMessage(Task)
_LESSON.fields_by_name['teacher_id'].has_options = True
_LESSON.fields_by_name['teacher_id']._options = _descriptor._ParseOptions(descriptor_pb2.FieldOptions(), _b('\020\001'))
_LESSON.fields_by_name['group_id'].has_options = True
_LESSON.fields_by_name['group_id']._options = _descriptor._ParseOptions(descriptor_pb2.FieldOptions(), _b('\020\001'))
# @@protoc_insertion_point(module_scope)

View File

View File

View File

@@ -0,0 +1,55 @@
from zvk.bot.bot import Bot
from zvk.plugins.vk.command import command, Argument
from zvk.plugins.vk.command_parser import Echo
from zvk.util import emoji
@command('list_plugins',
permissions=['admin'])
async def command_list_plugins(bot: Bot, echo: Echo):
plugin_lines = [f'- {plugin} {emoji.CHECK if plugin.is_activated else emoji.CROSS}' for plugin in bot.plugins.values()]
plugin_str = '\n'.join(sorted(plugin_lines))
s = f'{len(bot.plugins)} plugins:\n{plugin_str}'
await echo(s)
async def switch_plugin(bot: Bot, echo: Echo, plugin_name: str, new_state: bool):
if plugin_name not in bot.plugins:
await echo('No such plugin')
return
plugin = bot.plugins[plugin_name]
if plugin.is_activated == new_state:
await echo(f'Already {new_state}')
return
if new_state:
plugin.activate(bot.event_queue)
else:
plugin.deactivate(bot.event_queue)
await echo('Ok')
@command('disable_plugin',
Argument('plugin_name', type=str),
permissions=['admin'])
async def command_disable_plugin(bot: Bot, echo: Echo, plugin_name: str):
await switch_plugin(bot, echo, plugin_name, False)
@command('enable_plugin',
Argument('plugin_name', type=str),
permissions=['admin'])
async def command_enable_plugin(bot: Bot, echo: Echo, plugin_name: str):
await switch_plugin(bot, echo, plugin_name, True)
@command('die',
permissions=['admin'])
async def command_die(bot: Bot, echo: Echo):
await echo(emoji.SKULL)
bot.die()

View File

@@ -0,0 +1,97 @@
import asyncio
from datetime import datetime
from typing import List
from zvk.plugins.vk.api import VKApi
from zvk.plugins.vk.command import command
from zvk.plugins.vk.command_parser import Echo
from zvk.plugins.vk.upload import upload_image
from zvk.util.db import Database
from zvk.util.download import download_file
from zvk.util.network import Network
from zvk.util.zlogging import logger
GI_IMAGES_TO_ATTACH = 10
async def search_google_images(net: Network,
q: str,
start: int) -> List[str]:
params = {
'q': q,
'searchType': 'image',
'cx': '000676372658842926074:vgd-eou7wlq',
'key': 'AIzaSyBO15j9tuSmLhxhXMQyKGL2YzHnFhBP8o4',
'start': start
}
logger.info(f'Calling google api with params {params}')
_, data = await net.get_json('https://www.googleapis.com/customsearch/v1', params=params)
items = data.get('items', [])
results = []
for item in items:
link = item.get('link', None)
if link:
results.append(link)
else:
logger.warning(f'Unexpected response format: {item}')
return results
async def download_and_upload_image(db: Database,
net: Network,
api: VKApi,
url: str) -> str:
try:
local_path = await download_file(db, net, url, target_directory='google_images')
return await upload_image(db, net, api, local_path)
except Exception as e:
logger.warning(f'Could not download/upload {url}: {e}')
# Almost no one is evil. Almost everything is broken.
@command('gi', whole_argstring=True)
async def command_gi(db: Database,
net: Network,
api: VKApi,
argstring: str,
echo: Echo):
started_at = datetime.utcnow()
logger.info(f'Searching google images for request {argstring}')
url_tasks = await asyncio.gather(search_google_images(net, argstring, start=1),
search_google_images(net, argstring, start=11))
urls = sum(url_tasks, [])
if len(urls) == 0:
await echo('No results O_o')
return
uploaded_objects = []
download_tasks = [download_and_upload_image(db, net, api, url) for url in urls]
for coro in asyncio.as_completed(download_tasks):
result = await coro
if result is None:
continue
uploaded_objects.append(result)
if len(uploaded_objects) == GI_IMAGES_TO_ATTACH:
break
logger.info(f'From {len(urls)} urls downloaded and uploaded {len(uploaded_objects)} images')
attachment = ','.join(uploaded_objects)
await echo(f'"{argstring}": '
f'{len(uploaded_objects)} images in '
f'{(datetime.utcnow() - started_at).total_seconds():.2f} secs',
attachment=attachment)

View File

@@ -0,0 +1,151 @@
import json
import zipfile
import datetime
from io import BytesIO
import pandas
from sqlalchemy import Column, Integer, String
from zvk.misc.timetable_pb2 import Timetable
from zvk.plugins.vk.api import VKApi
from zvk.plugins.vk.command import command, Argument
from zvk.plugins.vk.command_parser import Echo
from zvk.util.db import DBBase, Database
from zvk.util.network import Network
class TimetableJson(DBBase):
id = Column(Integer, primary_key=True)
json = Column(String)
@command('import_timetable',
permissions=['admin'])
async def command_import_timetable(attachments, echo: Echo, api: VKApi, net: Network, db: Database):
if 'attach1' not in attachments:
await echo('You should attach something')
return
doc_id = attachments['attach1']
info = await api.docs.getById(docs=doc_id)
url = info[0]['url']
_, timetable_zip_bytes = await net.get_bytes(url)
timetable_zip = zipfile.ZipFile(BytesIO(timetable_zip_bytes))
timetable_pb_bytes = timetable_zip.read('timetable.pb')
timetable = Timetable.FromString(timetable_pb_bytes)
term_start = datetime.datetime.fromtimestamp(timetable.properties.term_start / 1000)
weeks_count = timetable.properties.weeks_count
def convert_weeks(s):
if s == 'a':
return list(range(1, weeks_count + 1))
if s == 'o':
return list(range(1, weeks_count + 1, 2))
if s == 'e':
return list(range(2, weeks_count + 1, 2))
if s.startswith('c'):
return list(map(int, s[1:].split(',')))
raise Exception(f'Bad week identifier {s}')
timetable_dict = {
'term_start': term_start.timestamp(),
'weeks_count': weeks_count,
'lessons': []
}
for lesson in timetable.lesson:
timetable_dict['lessons'].append({
'day': lesson.day,
'time': [lesson.time[:4], lesson.time[4:]],
'weeks': convert_weeks(lesson.weeks),
'subject': timetable.subject[lesson.subject_id - 1].name,
'kind': timetable.kind[lesson.kind_id - 1].name,
'place': timetable.place[lesson.place_id - 1].name,
'teachers': [timetable.teacher[teacher_id - 1].name
for teacher_id in lesson.teacher_id]
})
timetable_json = json.dumps(timetable_dict)
with db as session:
session.query(TimetableJson).delete()
timetable_json = TimetableJson(json=timetable_json)
session.add(timetable_json)
await echo(f'Imported timetable: {len(timetable_dict["lessons"])} lessons')
def get_day_timetable(timetable, now):
dt_from_start = now - datetime.datetime.utcfromtimestamp(timetable['term_start'])
days_from_start = dt_from_start // datetime.timedelta(days=1)
week_number = days_from_start // 7 + 1
day = days_from_start % 7 + 1
today = []
for lesson in timetable['lessons']:
if lesson['day'] != day or week_number not in lesson['weeks']:
continue
today.append(lesson)
today = sorted(today, key=lambda x: x['time'][0])
return today, week_number
def format_lesson(lesson):
ftime = lambda s: f'{s[:2]}:{s[2:]}'
return f'{ftime(lesson["time"][0])}->{ftime(lesson["time"][1])} ' \
f'{lesson["kind"]} {lesson["subject"]} {lesson["place"]}'
@command('tt',
Argument('n', type=int, default=2, nargs='?'))
async def command_tt(echo: Echo, db: Database, n):
if n <= 0 or n > 10:
await echo('fuk you')
return
with db as session:
timetable_json = session.query(TimetableJson).first().json
timetable = json.loads(timetable_json)
day_names = {
-1: 'Вчера',
0: 'Сегодня',
1: 'Завтра',
2: 'Послезавтра',
}
printed_nothing = True
dayds = list(range(min(0, n), max(1, n)))
for dayd in dayds:
dt = datetime.datetime.utcnow() + datetime.timedelta(days=dayd)
rounded = pandas.Timestamp(dt).floor('d').to_pydatetime()
tt, week_number = get_day_timetable(timetable, dt)
fmt = '\n'.join(map(format_lesson, tt))
if not fmt:
continue
day_name = day_names.get(dayd, f'{rounded.strftime("%b %d - %a")}')
await echo(f'{day_name} Неделя #{week_number}:\n{fmt}')
printed_nothing = False
if printed_nothing:
await echo(f'ниче нету')

View File

View File

@@ -0,0 +1,16 @@
from zvk.bot.trunk import Trunk
from zvk.event.consumer import on_startup
from zvk.plugins.vk.api import VKApi
from zvk.util.zlogging import logger
@on_startup
async def who_am_i(api: VKApi, trunk: Trunk):
owner = (await api.users.get())[0]
owner_id = owner['id']
trunk.set('owner_id', owner_id)
owner_name = f'{owner["first_name"]} {owner["last_name"]}'
trunk.set('owner_name', owner_name)
logger.info(f'Running @{owner_id} as {owner_name}')

View File

@@ -0,0 +1,41 @@
from typing import List, Set, Dict
from zvk.bot.trunk import Trunk
from zvk.event.consumer import on_startup
from zvk.util.zlogging import logger
class PermissionManager:
categories: Dict[str, Set[int]]
def __init__(self, config: dict):
self.categories = dict()
for category, members in config['permissions'].items():
members = set(members)
self.categories[category] = members
def update_owner(self, owner_id):
for members in self.categories.values():
if 0 in members:
members.remove(0)
members.add(owner_id)
def shall_pass(self, user_id: int, permissions: List[str]):
for permission_category in permissions:
if user_id in self.categories.get(permission_category, set()):
return True
return False
@on_startup
async def initialize_permissions(config: dict, trunk: Trunk, bot):
permission_manager = PermissionManager(config)
permission_manager.update_owner(await trunk.get('owner_id'))
trunk.set('permissions', permission_manager)
logger.info(f'{len(permission_manager.categories)} user categories initialized')

View File

View File

@@ -0,0 +1,9 @@
from zvk.event.periodic import periodic
from zvk.plugins.vk.api import VKApi
from zvk.util.zlogging import logger
@periodic(period_secs=300)
async def periodic_mark_online(api: VKApi):
await api.account.setOnline(voip=0)
logger.info('Marked as online')

View File

@@ -0,0 +1,85 @@
import asyncio
from datetime import datetime, timedelta
from sqlalchemy import DateTime, Integer
from zvk.bot.trunk import Trunk
from zvk.event.consumer import event_consumer
from zvk.plugins.vk.api import VKApi
from zvk.plugins.vk.event_type import VKEventType
from zvk.plugins.vk.longpoll import LongpollEvent
from zvk.util.db import DBBase, NNColumn, Database
from zvk.util.zlogging import logger
ACTIVITY_BREAK_TIMEOUT = timedelta(minutes=5)
class UserActivityPeriod(DBBase):
id = NNColumn(Integer, primary_key=True)
user_id = NNColumn(Integer)
period_start = NNColumn(DateTime)
period_end = NNColumn(DateTime)
@event_consumer(consumes=[LongpollEvent.FIRST_EVENT])
async def online_init(api: VKApi, trunk: Trunk, db: Database):
online_users = set()
r = await api.friends.get()
friend_list = r['items']
now = datetime.utcnow()
r = await api.users.get(user_ids=','.join(map(str, friend_list)), fields='online')
with db as session:
for i in r:
if i['online'] == 1:
online_users.add(i['id'])
session.add(UserActivityPeriod(
user_id=i['id'],
period_start=now,
period_end=now,
))
logger.info(f'Initialized online tracker: {len(online_users)}')
trunk.set('online_users', online_users)
@event_consumer(consumes=[VKEventType.USER_CAME_OFFLINE, VKEventType.USER_CAME_ONLINE])
async def online_listener(trunk: Trunk, db: Database, event_type, vk_event_args):
online_users = await trunk.get('online_users')
if event_type == VKEventType.USER_CAME_OFFLINE:
pass
if event_type == VKEventType.USER_CAME_ONLINE:
pass
user_id = -vk_event_args[0]
timestamp = vk_event_args[2]
with db as session:
if event_type == VKEventType.USER_CAME_OFFLINE and user_id in online_users:
last_period = session \
.query(UserActivityPeriod) \
.filter_by(user_id=user_id) \
.order_by(UserActivityPeriod.id.desc()) \
.first()
if not last_period:
return
last_period.period_end = datetime.fromtimestamp(timestamp)
online_users.remove(user_id)
if event_type == VKEventType.USER_CAME_ONLINE and user_id not in online_users:
session.add(UserActivityPeriod(
user_id=user_id,
period_start=datetime.fromtimestamp(timestamp),
period_end=datetime.fromtimestamp(timestamp),
))
online_users.add(user_id)

View File

86
src/zvk/plugins/vk/api.py Normal file
View File

@@ -0,0 +1,86 @@
from __future__ import annotations
import asyncio
from zvk.util.network import Network
from zvk.util.zlogging import logger
# BOT_MESSAGE_RANDOM_ID_MIN = 1337000000
# BOT_MESSAGE_RANDOM_ID_MAX = 1338000000
API_VERSION = '5.85'
class MagicAccumulatingAttributeCatcher:
api: VKApi
full_method_name: str
def __init__(self, api, full_method_name=None):
self.api = api
if full_method_name is None:
full_method_name = ''
self.full_method_name = full_method_name
def __getattr__(self, item):
return MagicAccumulatingAttributeCatcher(self.api, f'{self.full_method_name}.{item}')
async def __call__(self, **kwargs):
return await self.api.call_method(self.full_method_name, **kwargs)
API_CALL_RETRY_COUNT = 3
class VKApi:
"""
VK api asynchronous interaction interface. Supports dope syntax like
`await api.messages.get()`.
"""
_access_token: str
_net: Network
def __init__(self, config, net):
self._access_token = config['api']['access_token']
self._net = net
def __getattr__(self, name):
return MagicAccumulatingAttributeCatcher(self, name)
async def call_method(self, full_method_name, **params):
url = f'https://api.vk.com/method/{full_method_name}'
params['access_token'] = self._access_token
params['v'] = API_VERSION
# TODO: random_id management
# if self.name == 'messages.send' and 'random_id' not in params:
# params['random_id'] = random.randint(BOT_MESSAGE_RANDOM_ID_MIN, BOT_MESSAGE_RANDOM_ID_MAX)
# if self.name == 'messages.send' and 'message' in params:
# params['message'] = params['message'].replace('--', '--')
# filter empty values
params = {k: v for k, v in params.items() if v is not None}
for retry in range(API_CALL_RETRY_COUNT):
response, result = await self._net.post_json(url, sequential=True, data=params)
if 'error' in result:
error = result['error']
if error['error_code'] == 6:
logger.warning('Too many requests, waiting and retrying...')
await asyncio.sleep(1)
continue
raise RuntimeError(f'A significant VKApi error occurred {full_method_name}({params}) -> {error}')
if 'response' not in result:
raise RuntimeError(f'Malformed api response {full_method_name}({params}) -> {result}')
logger.debug(f'VKApi call {full_method_name}({params}) -> {result}')
return result['response']
raise RuntimeError(f'VKApi call unsuccessful after retries: {full_method_name}({params})')

View File

@@ -0,0 +1,120 @@
import argparse
import shlex
from typing import List
from zvk.event.consumer import EventConsumer
from zvk.event.event import Event
from zvk.plugins.vk.command_parser import CommandEventType
from zvk.util.zlogging import logger
class CommandParseException(Exception):
def __init__(self, message):
self.message = message
class CommandArgumentParser(argparse.ArgumentParser):
def print_help(self, file=None):
pass
def exit(self, status=0, message=None):
raise CommandParseException(self.format_help())
def error(self, message):
raise CommandParseException(f'Command .{self.prog} {message}')
class Argument:
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
class CommandEventConsumer(EventConsumer):
command_name: str
whole_argstring: bool
parser: CommandArgumentParser
allowed_permission_categories: List[str]
def __init__(self,
command_name: str,
*args: Argument,
whole_argstring: bool = False,
description: str = None,
permissions: List[str] = None):
self.command_name = command_name
self.whole_argstring = whole_argstring
self.parser = CommandArgumentParser(prog=self.command_name)
self.allowed_permission_categories = permissions
if description is None:
self.parser.description = 'TODO'
else:
self.parser.description = description
if self.whole_argstring:
if args:
raise TypeError('Whole argstring cannot have other args')
self.parser.add_argument('argstring',
type=str,
nargs=argparse.REMAINDER,
metavar='...',
help='the whole command line')
for arg in args:
self.parser.add_argument(*arg.args, **arg.kwargs)
super().__init__(consumes=[CommandEventType(command_name=command_name)])
def parse_argstring(self, argstring: str) -> dict:
if self.whole_argstring:
return dict(argstring=argstring)
try:
parts = shlex.split(argstring)
except ValueError as e:
raise CommandParseException(e.args[0])
namespace = self.parser.parse_args(args=parts)
return vars(namespace)
async def consume(self, event: Event):
if self.allowed_permission_categories is not None:
from_id = event.env['message'].from_id
permission_manager = await event.env['trunk'].get('permissions')
if not permission_manager.shall_pass(from_id, self.allowed_permission_categories):
# TODO: fancier
await event.env['echo']('Access denied')
logger.warning(f'Access denied {from_id} {event}')
return
try:
event.env.update(self.parse_argstring(event.env['command_argstring']))
except CommandParseException as e:
await event.env['echo'](e.message)
logger.warning(f'Wrong command {event} {e.message}')
return
async for output_event in super().consume(event):
yield output_event
def command(command_name: str,
*args: Argument,
whole_argstring: bool = False,
description: str = None,
permissions: List[str] = None) -> EventConsumer:
if callable(command_name):
# direct decoration
raise TypeError('Direct command decoration is forbidden')
return CommandEventConsumer(
command_name,
*args,
whole_argstring=whole_argstring,
description=description,
permissions=permissions)

View File

@@ -0,0 +1,66 @@
import html
import json
import re
from dataclasses import dataclass
from zvk.event.consumer import event_consumer
from zvk.event.event import Event
from zvk.plugins.vk.api import VKApi
from zvk.plugins.vk.event_type import ParsedEventType
from zvk.plugins.vk.message_parser import Message
from zvk.util import emoji
from zvk.util.zlogging import logger
# old
# COMMAND_REGEX = r'\.(\S+)\s*(.*)'
# starts with comma to avoid collisions with old schema
COMMAND_REGEX = r'\,(\S+)\s*(.*)'
def preprocess_argstring(s):
return html.unescape(s) \
.replace('', '--') \
.replace('«', '<<') \
.replace('»', '>>') \
.replace('<br>', '\n')
@dataclass
class CommandEventType:
command_name: str
def __hash__(self):
return hash(self.command_name)
@dataclass
class Echo:
api: VKApi
peer_id: int
async def __call__(self, message='', notext=False, **kwargs):
if not notext:
result = await self.api.messages.send(peer_id=self.peer_id, message=f'{emoji.ROBOT}: {message}', **kwargs)
else:
result = await self.api.messages.send(peer_id=self.peer_id, **kwargs)
logger.info(f'Echo to {self.peer_id}: {message}{f"extra: {kwargs}" if kwargs else ""}')
return result
@event_consumer(consumes=[ParsedEventType.MESSAGE])
async def command_parser(api: VKApi, message: Message):
match = re.fullmatch(COMMAND_REGEX, message.text)
if match is None:
return
command_name = match.group(1)
command_argstring = preprocess_argstring(match.group(2))
yield Event(CommandEventType(command_name),
command_name=command_name,
command_argstring=command_argstring,
echo=Echo(api, message.peer_id),
attachments=json.loads(message.attachments_json))

View File

@@ -0,0 +1,25 @@
import json
from sqlalchemy import Integer, String
from zvk.event.consumer import event_consumer
from zvk.util.db import DBBase, NNColumn, Database
from zvk.util.zlogging import logger
from zvk.plugins.vk.event_type import VKEventType
class VKEvent(DBBase):
id = NNColumn(Integer, primary_key=True)
vk_event_type_id = NNColumn(Integer)
vk_event_args_json = NNColumn(String)
@event_consumer(consumes=list(VKEventType))
async def save_event(db: Database, event_type: VKEventType, vk_event_args):
logger.debug(f'Persisting event {event_type} {vk_event_args}')
with db as session:
session.add(VKEvent(
vk_event_type_id=event_type.value,
vk_event_args_json=json.dumps(vk_event_args)))

View File

@@ -0,0 +1,36 @@
from enum import Enum, auto
class VKEventType(Enum):
MESSAGE_FLAG_REPLACEMENT = 1
MESSAGE_FLAG_SETTING = 2
MESSAGE_FLAG_REMOVAL = 3
MESSAGE_NEW = 4
MESSAGE_EDIT = 5
MESSAGE_READ_INCOMING = 6
MESSAGE_READ_OUTGOING = 7
USER_CAME_ONLINE = 8
USER_CAME_OFFLINE = 9
CHAT_FLAG_REPLACEMENT = 10
CHAT_FLAG_SETTING = 11
CHAT_FLAG_REMOVAL = 12
MESSAGE_DELETED = 13
MESSAGE_RESTORED = 14
USER_TYPING = 61
USER_TYPING_CHAT = 62
CALL_NEW = 70
UNREAD_COUNTER_UPDATE = 80
CHAT_NOTIFICATION_SETTINGS_CHANGE = 114
class ParsedEventType(Enum):
MESSAGE = auto()
MESSAGE_EDIT = auto()

View File

@@ -0,0 +1,94 @@
from __future__ import annotations
import asyncio
from enum import Enum, auto
from zvk.bot.event_type import BotEventType
from zvk.util.network import Network
from zvk.event.consumer import event_consumer
from zvk.event.event import Event
from zvk.util.zlogging import logger
from zvk.plugins.vk.api import VKApi
from zvk.plugins.vk.event_type import VKEventType
# attachments + extended + extra + random_id
LONGPOLL_MODE = 2 | 8 | 64 | 128
LONGPOLL_VERSION = 3
CALLS_PER_SERVER = 100
# TODO: increase?
WAIT_SECONDS = 25
class LongpollEvent(Enum):
FIRST_EVENT = auto()
@event_consumer(consumes=[BotEventType.STARTUP])
async def longpoll_loop(bot, net: Network, api: VKApi):
last_event_timestamp = None
sent_announce = False
while True:
try:
server = await api.messages.getLongPollServer()
if last_event_timestamp is None:
last_event_timestamp = server['ts']
logger.info(f'Starting longpolling at {last_event_timestamp} from {server["server"]}')
for longpoll_call in range(CALLS_PER_SERVER):
poll_url = f'https://{server["server"]}'
poll_params = {
'act': 'a_check',
'key': server['key'],
'ts': last_event_timestamp,
'wait': WAIT_SECONDS,
'mode': LONGPOLL_MODE,
'version': LONGPOLL_VERSION
}
response, payload = await net.get_json(poll_url, params=poll_params)
failed_code = payload.get('failed', 0)
if failed_code > 0:
logger.warn(f'Longpolling call failed with {payload}')
if failed_code == 1:
# ts is lost
last_event_timestamp = payload['ts']
continue
elif failed_code == 2:
# key is stale
break
elif failed_code == 3:
# user is lost
break
elif failed_code == 4:
# bad version
logger.error(f'Bad LP version {LONGPOLL_VERSION}')
bot.die()
logger.debug(f'Longpoll tick @{last_event_timestamp} {len(payload["updates"])} events')
last_event_timestamp = payload['ts']
for update in payload['updates']:
vk_event_type_id = update[0]
vk_event_args = update[1:]
vk_event_type = VKEventType(vk_event_type_id)
logger.info(f'{vk_event_type} {vk_event_args}')
if not sent_announce:
yield Event(LongpollEvent.FIRST_EVENT)
sent_announce = True
yield Event(vk_event_type, vk_event_args=vk_event_args)
except asyncio.CancelledError:
raise
except Exception:
logger.exception(f'Longpoll failed o_O')
last_event_timestamp = None

View File

@@ -0,0 +1,26 @@
from __future__ import annotations
from enum import Enum
from typing import Set
import numpy
class MessageFlag(Enum):
UNREAD = 1 << 0
OUTBOX = 1 << 1
REPLIED = 1 << 2
IMPORTANT = 1 << 3
FRIENDS = 1 << 5
SPAM = 1 << 6
DELETED = 1 << 7
DELETED_FOR_ALL = 1 << 17
@staticmethod
def parse_flags(flags: int) -> Set[MessageFlag]:
return {i for i in MessageFlag if i.value & flags > 0}
@staticmethod
def encode_flags(flags: Set[MessageFlag]) -> int:
# noinspection PyTypeChecker
return numpy.bitwise_or.reduce([i.value for i in flags])

View File

@@ -0,0 +1,130 @@
import json
from datetime import datetime
from sqlalchemy import Boolean, DateTime, Integer, String
from zvk.bot.trunk import Trunk
from zvk.event.consumer import event_consumer
from zvk.event.event import Event
from zvk.plugins.vk.event_type import ParsedEventType, VKEventType
from zvk.plugins.vk.message_flags import MessageFlag
from zvk.util.db import DBBase, Database, NNColumn
class Message(DBBase):
id = NNColumn(Integer, primary_key=True)
message_id = NNColumn(Integer, nullable=False, unique=True)
from_id = NNColumn(Integer)
to_id = NNColumn(Integer)
peer_id = NNColumn(Integer)
flags = NNColumn(Integer)
timestamp = NNColumn(DateTime)
text = NNColumn(String)
extra_fields_json = NNColumn(String)
attachments_json = NNColumn(String)
random_id = NNColumn(Integer)
is_outgoing = NNColumn(Boolean)
is_bot_message = NNColumn(Boolean)
@event_consumer(consumes=[VKEventType.MESSAGE_NEW])
async def new_message(db: Database, trunk: Trunk, vk_event_args):
message_id, flags, peer_id, timestamp, text, extra_fields, attachments, random_id = vk_event_args
peer_id = peer_id
if 'from' in extra_fields:
# message to a chat
from_id = extra_fields['from']
to_id = peer_id
else:
# direct message
from_id = peer_id
to_id = await trunk.get('owner_id')
flag_set = MessageFlag.parse_flags(flags)
is_outgoing = MessageFlag.OUTBOX in flag_set
is_bot_message = False
# TODO: api.messages.getById
# if 'fwd_count' in extra_fields:
# pass
# TODO: also trigger on message edit
with db as session:
message = Message(
message_id=message_id,
from_id=from_id,
to_id=to_id,
peer_id=peer_id,
flags=flags,
timestamp=datetime.fromtimestamp(timestamp),
text=text,
extra_fields_json=json.dumps(extra_fields),
attachments_json=json.dumps(attachments),
random_id=random_id,
is_outgoing=is_outgoing,
is_bot_message=is_bot_message,
)
session.add(message)
yield Event(ParsedEventType.MESSAGE, message=message)
# @event_consumer(consumes=[VKEventType.MESSAGE_EDIT])
# async def edit_message(db: Database, trunk: Trunk, vk_event_args):
# message_id, mask, peer_id, timestamp, new_text, attachments, _ = vk_event_args
#
#
#
# flag_set = MessageFlag.parse_flags(flags)
#
# is_outgoing = MessageFlag.OUTBOX in flag_set
# is_bot_message = False
#
# # TODO: api.messages.getById
# # if 'fwd_count' in extra_fields:
# # pass
#
# # TODO: also trigger on message edit
#
# with db as session:
# message = Message(
# message_id=message_id,
#
# from_id=from_id,
# to_id=to_id,
#
# flags=flags,
# timestamp=datetime.fromtimestamp(timestamp),
#
# text=text,
#
# extra_fields_json=json.dumps(extra_fields),
# attachments_json=json.dumps(attachments),
#
# random_id=random_id,
#
# is_outgoing=is_outgoing,
# is_bot_message=is_bot_message,
# )
#
# session.add(message)
#
# yield Event(ParsedEventType.MESSAGE, message=message)

View File

@@ -0,0 +1,65 @@
import mimetypes
from sqlalchemy import Integer, String
from zvk.plugins.vk.api import VKApi
from zvk.util.db import DBBase, NNColumn, Database
from zvk.util.network import Network
from zvk.util.zlogging import logger
class CachedUpload(DBBase):
id = NNColumn(Integer, primary_key=True)
local_path = NNColumn(String)
vk_object = NNColumn(String)
def is_local_file_an_image(local_path):
mimetype, _ = mimetypes.guess_type(local_path)
return mimetype in ['image/jpeg', 'image/png', 'image/gif']
async def upload_image(db: Database,
net: Network,
api: VKApi,
local_path: str,
use_cache: bool = True) -> str:
logger.info(f'Uploading local image {local_path}')
if use_cache:
with db as session:
cached_upload = session \
.query(CachedUpload) \
.filter_by(local_path=local_path) \
.first()
if cached_upload:
logger.info(f'Found {local_path} in upload cache: {cached_upload.vk_object}')
return cached_upload.vk_object
if not is_local_file_an_image(local_path):
raise ValueError(f'{local_path} does not look like an image!')
response = await api.photos.getMessagesUploadServer()
upload_url = response['upload_url']
with open(local_path, 'rb') as f:
files = {'photo': f}
_, response = await net.post_json(upload_url, data=files)
if 'server' not in response or 'photo' not in response or 'hash' not in response:
raise RuntimeError(f'Could not upload {local_path}')
logger.info(f'Uploaded {local_path} as {response}')
response = await api.photos.saveMessagesPhoto(**response)
if len(response) < 1:
raise RuntimeError(f'Could not save image {local_path}')
vk_object = f'photo{response[0]["owner_id"]}_{response[0]["id"]}'
if use_cache:
with db as session:
session.add(CachedUpload(local_path=local_path, vk_object=vk_object))
return vk_object

0
src/zvk/util/__init__.py Normal file
View File

59
src/zvk/util/db.py Normal file
View File

@@ -0,0 +1,59 @@
from functools import partial
from typing import Any, Callable
from sqlalchemy import create_engine, Column
from sqlalchemy.ext.declarative import declarative_base, DeclarativeMeta
from sqlalchemy.orm import sessionmaker
from sqlalchemy.orm.session import Session
class AutoTableNamer(DeclarativeMeta):
def __new__(cls, name, bases, classdict):
if '__tablename__' in classdict:
raise TypeError(f'Table name already defined for {name}')
classdict['__tablename__'] = f'{name}_auto'
res_cls = super().__new__(cls, name, bases, classdict)
res_cls.f = 1
return res_cls
class Database:
session: Session
def __init__(self, url, **kwargs):
self.session = None
self.engine = create_engine(url, **kwargs)
self._session_factory = sessionmaker(bind=self.engine, expire_on_commit=False)
def __enter__(self) -> Session:
# self.create_all()
if self.session is not None:
raise RuntimeError('nested db sessions')
self.session = self._session_factory()
return self.session
def __exit__(self, exc_type, exc_val, exc_tb):
if exc_type:
self.session.rollback()
else:
self.session.commit()
self.session.close()
self.session = None
def create_all(self):
DBBase.metadata.create_all(bind=self.engine)
DBBase: None = declarative_base(metaclass=AutoTableNamer)
NNColumn = partial(Column, nullable=False)

98
src/zvk/util/download.py Normal file
View File

@@ -0,0 +1,98 @@
import json
import mimetypes
import os
from slugify import slugify
from sqlalchemy import Integer, String
from zvk.util import paths
from zvk.util.db import DBBase, Database, NNColumn
from zvk.util.network import Network
from zvk.util.zlogging import logger
class CachedDownload(DBBase):
id = NNColumn(Integer, primary_key=True)
url = NNColumn(String)
params_json = NNColumn(String)
local_path = NNColumn(String)
async def download_file(db: Database,
net: Network,
url: str,
filename: str = None,
target_directory: str = 'unspecified',
params: dict = None,
use_cache: bool = True) -> str:
"""
Asynchronously downloads a file from an url and stores it locally.
"""
if filename is None:
filename = slugify(url)
if params is None:
params = {}
params_json = json.dumps(params)
logger.info(f'Downloading {url}({params}) to {target_directory}/{filename}')
if use_cache:
with db as session:
cached_download = session \
.query(CachedDownload) \
.filter_by(url=url,
params_json=params_json) \
.first()
if cached_download:
logger.info(f'Found {url}({params}) in download cache: {cached_download.local_path}')
if not os.path.exists(cached_download.local_path):
logger.warning(f'No file from {url}({params}) exists in {cached_download.local_path}, purging from cache')
session.delete(cached_download)
else:
return cached_download.local_path
response, content = await net.get_bytes(url, params=params)
if response.status != 200:
raise RuntimeError(f'Got a bad status code {response.status} from {url}({params})')
logger.info(f'Downloaded {len(content)} bytes from {url}({params})')
directory = os.path.join(paths.DOWNLOAD_DIR, target_directory)
os.makedirs(directory, exist_ok=True)
filename += guess_suffix(response)
local_path = os.path.join(directory, filename)
with open(local_path, 'wb') as file:
file.write(content)
if use_cache:
with db as session:
session.add(CachedDownload(url=url,
params_json=params_json,
local_path=local_path))
logger.info(f'Saved {len(content)} bytes from {url}({params}) to {local_path}')
return local_path
def guess_suffix(response):
mimetype = response.headers.get('Content-Type')
if not mimetype:
mimetype, _ = mimetypes.guess_type(str(response.real_url))
if mimetype == 'image/jpeg':
return f'.jpg'
elif mimetype == 'image/png':
return f'.png'
elif mimetype == 'image/gif':
return f'.gif'
return ''

14
src/zvk/util/emoji.py Normal file
View File

@@ -0,0 +1,14 @@
THUMBS_UP = '👍'
FLEX = '💪'
HEART = '💗'
RESTART = '🔄'
ROBOT = '🤖'
LEMON = '🍋'
PIE = '🍰'
OK = '👌'
HUNDRED = '💯'
COW_FACE = '🐮'
SKULL = '💀'
WINK = '😉'
CHECK = ''
CROSS = ''

107
src/zvk/util/network.py Normal file
View File

@@ -0,0 +1,107 @@
import asyncio
import json
from enum import Enum, auto
from typing import Any, Tuple
import aiohttp
from zvk.util.zlogging import logger
class ReturnType(Enum):
BYTES = auto()
TEXT = auto()
JSON = auto()
class Network:
"""
Encapsulation of asynchronous network interaction
"""
_timeout: int
_lock: asyncio.Lock
def __init__(self, config):
self._timeout = config['net']['timeout']
self._lock = None
def initialize(self):
self._lock = asyncio.Lock()
async def _request(self, method, url, sequential, return_type, **kwargs):
if sequential:
await self._lock.acquire()
result = None
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=self._timeout)) as session:
async with session.request(method, url, **kwargs) as response:
if return_type == ReturnType.BYTES:
result = await response.read()
if return_type == ReturnType.TEXT:
result = await response.text()
if return_type == ReturnType.JSON:
result = json.loads(await response.text())
if sequential:
self._lock.release()
return response, result
async def get_bytes(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, bytes]:
return await self._request(
method='GET',
url=url,
sequential=sequential,
return_type=ReturnType.BYTES,
**kwargs
)
async def post_bytes(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, bytes]:
return await self._request(
method='POST',
url=url,
sequential=sequential,
return_type=ReturnType.BYTES,
**kwargs
)
async def get_text(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, str]:
return await self._request(
method='GET',
url=url,
sequential=sequential,
return_type=ReturnType.TEXT,
**kwargs
)
async def post_text(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, str]:
return await self._request(
method='POST',
url=url,
sequential=sequential,
return_type=ReturnType.TEXT,
**kwargs
)
async def get_json(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, Any]:
return await self._request(
method='GET',
url=url,
sequential=sequential,
return_type=ReturnType.JSON,
**kwargs
)
async def post_json(self, url, sequential=False, **kwargs) -> Tuple[aiohttp.ClientResponse, Any]:
return await self._request(
method='POST',
url=url,
sequential=sequential,
return_type=ReturnType.JSON,
**kwargs
)

5
src/zvk/util/paths.py Normal file
View File

@@ -0,0 +1,5 @@
CONFIG_PATH = 'config.yaml'
PLUGIN_GLOB = 'src/zvk/plugins/**/*.py'
DOWNLOAD_DIR = 'downloads'

25
src/zvk/util/zlogging.py Normal file
View File

@@ -0,0 +1,25 @@
import logging
import sys
logger = logging.getLogger('zvk')
logger.setLevel(logging.INFO)
formatter = logging.Formatter(
'{asctime}.{msecs:03.0f} {levelname:>8} {module:>15.15}:{lineno:03d} - {message}',
style='{',
datefmt='%Y-%m-%d %H:%M:%S')
# debug_file = logging.FileHandler('debug.log')
# debug_file.setLevel(logging.DEBUG)
# debug_file.setFormatter(formatter)
# logger.addHandler(debug_file)
#
# warning_file = logging.FileHandler('warning.log')
# warning_file.setLevel(logging.WARNING)
# warning_file.setFormatter(formatter)
# logger.addHandler(warning_file)
info_console = logging.StreamHandler(stream=sys.stderr)
# info_console.setLevel(logging.INFO)
info_console.setFormatter(formatter)
logger.addHandler(info_console)