mirror of
https://github.com/wassname/catalyst.git
synced 2026-07-25 13:10:33 +08:00
test.py is successfully relaying all datasource events.
This commit is contained in:
+67
-24
@@ -1,67 +1,110 @@
|
||||
|
||||
from data.sources.equity import *
|
||||
import time
|
||||
import logging
|
||||
|
||||
class DataFeed(object):
|
||||
|
||||
def __init__(self, db, logger):
|
||||
self.logger = logger
|
||||
def __init__(self, db, subscriber_count):
|
||||
self.logger = logging.getLogger()
|
||||
self.db = db
|
||||
self.data_workers = {}
|
||||
self.data_address = "tcp://127.0.0.1:{port}".format(port=10101)
|
||||
self.sync_address = "tcp://127.0.0.1:{port}".format(port=10102)
|
||||
self.feed_address = "tcp://127.0.0.1:{port}".format(port=10103)
|
||||
self.data_buffer = {} #source_id -> []
|
||||
self.subscriber_count = subscriber_count
|
||||
|
||||
def start_data_workers(self):
|
||||
"""Start a sub-process for each datasource."""
|
||||
|
||||
# Socket to receive signals
|
||||
syncservice = self.context.socket(zmq.REP)
|
||||
syncservice.bind(self.sync_address)
|
||||
|
||||
|
||||
emt1 = EquityMinuteTrades(133, self.db, self.data_address, self.sync_address, 1, self.logger)
|
||||
emt1 = EquityMinuteTrades(133, self.db, self.data_address, self.sync_address, 1)
|
||||
self.data_workers[1] = emt1
|
||||
self.data_buffer[1] = []
|
||||
emt1.start()
|
||||
emt2 = EquityMinuteTrades(134, self.db, self.data_address, self.sync_address, 2, self.logger)
|
||||
emt2 = EquityMinuteTrades(134, self.db, self.data_address, self.sync_address, 2)
|
||||
self.data_workers[2] = emt1
|
||||
self.data_buffer[2] = []
|
||||
emt2.start()
|
||||
|
||||
workers = 0
|
||||
while workers < len(self.data_workers):
|
||||
self.logger.info("ds processes launched")
|
||||
|
||||
def sync_clients(self):
|
||||
# Socket to receive signals
|
||||
self.logger.info("waiting for all datasources and clients to be ready")
|
||||
self.syncservice = self.context.socket(zmq.REP)
|
||||
self.syncservice.bind(self.sync_address)
|
||||
|
||||
subscribers = 0
|
||||
while subscribers < (self.subscriber_count + len(self.data_workers)):
|
||||
self.logger.info("sync'ing {count}".format(count=subscribers))
|
||||
# wait for synchronization request
|
||||
msg = syncservice.recv()
|
||||
msg = self.syncservice.recv()
|
||||
# send synchronization reply
|
||||
syncservice.send('')
|
||||
workers += 1
|
||||
self.syncservice.send('')
|
||||
subscribers += 1
|
||||
|
||||
syncservice.close()
|
||||
|
||||
self.logger.info("{count} ds processes launched".format(count=workers))
|
||||
self.syncservice.close()
|
||||
self.logger.info("sync'd all datasources and clients")
|
||||
|
||||
def run(self):
|
||||
# Prepare our context and sockets
|
||||
self.context = zmq.Context()
|
||||
|
||||
counter = 0
|
||||
self.start_data_workers()
|
||||
|
||||
#create the data sink. Based on http://zguide.zeromq.org/py:tasksink2
|
||||
#see: http://zguide.zeromq.org/py:taskwork2
|
||||
self.data_socket = self.context.socket(zmq.PULL)
|
||||
self.data_socket.bind(self.data_address)
|
||||
|
||||
counter = 0
|
||||
#create the feed
|
||||
self.feed_socket = self.context.socket(zmq.PUSH)
|
||||
self.feed_socket.bind(self.feed_address)
|
||||
|
||||
self.start_data_workers()
|
||||
#wait for all feed subscribers
|
||||
self.sync_clients()
|
||||
|
||||
self.logger.info("entering feed loop on {addr}".format(addr=self.data_address))
|
||||
|
||||
while True:
|
||||
message = self.data_socket.recv()
|
||||
event = json.loads(message)
|
||||
#self.logger.info(message)
|
||||
counter = counter + 1
|
||||
if(event['type'] == "DONE"):
|
||||
source = event['s']
|
||||
#self.logger.info(" count " + str(counter) + " - " + str(event['dt']))
|
||||
if(event["type"] == "DONE"):
|
||||
self.logger.info("DONE")
|
||||
source = event[u's']
|
||||
if(self.data_workers.has_key(source)):
|
||||
del(self.data_workers[source])
|
||||
if(len(self.data_workers) == 0):
|
||||
break
|
||||
else:
|
||||
self.data_buffer[event[u's']].append(event)
|
||||
counter = counter + 1
|
||||
self.send_earliest_event()
|
||||
|
||||
|
||||
#drain any remaining messages in the buffer
|
||||
self.send_earliest_event(drain=True)
|
||||
|
||||
self.logger.info("Collected {n} messages".format(n=counter))
|
||||
self.data_socket.close()
|
||||
self.context.term()
|
||||
self.feed_socket.close()
|
||||
self.context.term()
|
||||
|
||||
def send_earliest_event(self, drain=False):
|
||||
earliest = None
|
||||
next_source = None
|
||||
while True: #send messages as long as we have >0 messages from each source
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if(not drain and len(events) == 0 and self.data_workers.has_key(source)):
|
||||
#there's no way to know that we have the next message
|
||||
return
|
||||
if(len(events) > 0 and (earliest == None or earliest > events[0])):
|
||||
earliest = events[0]['dt']
|
||||
next_source = source
|
||||
|
||||
|
||||
event = self.data_buffer[next_source].pop(0)
|
||||
self.feed_socket.send(json.dumps(event))
|
||||
+14
-14
@@ -6,6 +6,7 @@ import json
|
||||
import pytz
|
||||
import copy
|
||||
import multiprocessing
|
||||
import logging
|
||||
from pymongo import ASCENDING, DESCENDING
|
||||
|
||||
from backtest.util import *
|
||||
@@ -13,13 +14,13 @@ from backtest.util import *
|
||||
|
||||
class EquityMinuteTrades(object):
|
||||
|
||||
def __init__(self, sid, db, data_address, sync_address, source_id, logger):
|
||||
def __init__(self, sid, db, data_address, sync_address, source_id):
|
||||
self.sid = sid
|
||||
self.db = db
|
||||
self.source_id = source_id
|
||||
self.logger = logger
|
||||
self.data_address = data_address
|
||||
self.logger = logging.getLogger()
|
||||
self.sync_address = sync_address
|
||||
self.data_address = data_address
|
||||
self.logger.info("data address is {ds}".format(ds=data_address))
|
||||
|
||||
self.cur_event = None
|
||||
@@ -29,9 +30,18 @@ class EquityMinuteTrades(object):
|
||||
self.proc.start()
|
||||
|
||||
def run(self):
|
||||
self.logger.info("starting data source:{sid}".format(sid=self.sid))
|
||||
self.logger.info("starting data source:{sid} on {addr}".format(sid=self.sid, addr=self.data_address))
|
||||
self.context = zmq.Context()
|
||||
|
||||
#synchronize with feed
|
||||
sync_socket = self.context.socket(zmq.REQ)
|
||||
sync_socket.connect(self.sync_address)
|
||||
# send a synchronization request to the feed
|
||||
sync_socket.send('')
|
||||
# wait for synchronization reply from the feed
|
||||
sync_socket.recv()
|
||||
sync_socket.close()
|
||||
|
||||
#create the data sink. Based on http://zguide.zeromq.org/py:tasksink2
|
||||
self.data_socket = self.context.socket(zmq.PUSH)
|
||||
self.data_socket.connect(self.data_address)
|
||||
@@ -42,16 +52,6 @@ class EquityMinuteTrades(object):
|
||||
slave_ok=True)
|
||||
self.logger.info("found {count} events".format(count=eventQS.count()))
|
||||
|
||||
#synchronize with feed
|
||||
syncclient = self.context.socket(zmq.REQ)
|
||||
syncclient.connect(self.sync_address)
|
||||
|
||||
# send a synchronization request
|
||||
syncclient.send('')
|
||||
# wait for synchronization reply
|
||||
syncclient.recv()
|
||||
|
||||
syncclient.close()
|
||||
|
||||
for doc in eventQS:
|
||||
doc_dt = doc['dt'].replace(tzinfo = pytz.utc)
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
import zmq
|
||||
import logging
|
||||
import datetime
|
||||
import json
|
||||
import config
|
||||
import multiprocessing
|
||||
|
||||
class Transform(object):
|
||||
"""Parent class for feed transforms. Subclass to create a new derived value from the combined feed."""
|
||||
|
||||
def __init__(self, feed_address, result_address, sync_address, config_dict):
|
||||
"""
|
||||
feed_address - zmq socket address, Transform will CONNECT a PULL socket and receive messages until "DONE" is received.
|
||||
result_address - zmq socket address, Transform will CONNECT a PUSH socket and send messaes until feed_socket receives "DONE"
|
||||
sync_address - zmq socket address, Transform will CONNECT a REQ socket and send/receive one message before entering feed loop
|
||||
config - must be a config.Config object with at least an entry for 'name':string value
|
||||
"""
|
||||
self.logger = logging.getLogger()
|
||||
self.feed_address = feed_address
|
||||
self.result_address = result_address
|
||||
self.sync_address = sync_address
|
||||
self.config = config.Config(config_dict)
|
||||
self.name = self.config.get_string('name')
|
||||
self.state = {}
|
||||
self.state['name'] = self.name
|
||||
|
||||
def run(self):
|
||||
self.context = zmq.Context()
|
||||
|
||||
self.logger.info("starting {name} transform".format(name = self.name))
|
||||
#create the feed PULL.
|
||||
self.feed_socket = self.context.socket(zmq.PULL)
|
||||
self.feed_socket.connect(self.feed_address)
|
||||
|
||||
#create the result PUSH
|
||||
self.result_socket = self.context.socket(zmq.PUSH)
|
||||
self.result_socket.connect(self.result_address)
|
||||
|
||||
self.logger.info("sync'ing feed from {name}".format(name = self.name))
|
||||
#synchronize with feed
|
||||
sync_socket = self.context.socket(zmq.REQ)
|
||||
sync_socket.connect(self.sync_address)
|
||||
# send a synchronization request to the feed
|
||||
sync_socket.send('')
|
||||
# wait for synchronization reply from the feed
|
||||
sync_socket.recv()
|
||||
sync_socket.close()
|
||||
|
||||
self.logger.info("starting {name} event loop".format(name = self.name))
|
||||
|
||||
while True:
|
||||
message = self.feed_socket.recv()
|
||||
self.logger.info("got feed message at {name}".format(name=self.name))
|
||||
if(message == "DONE"):
|
||||
break;
|
||||
event = json.loads(message)
|
||||
cur_state = update(event)
|
||||
|
||||
self.result_socket.send(json.dumps(cur_state))
|
||||
self.logger.info("sent message from {name}".format(name=self.name))
|
||||
|
||||
def update(self, event):
|
||||
return {}
|
||||
|
||||
class Merge(Transform):
|
||||
""" Merge data feed and array of transform feeds into a single result vector.
|
||||
PULL from feed
|
||||
PULL from child transforms
|
||||
PUSH to client
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, feed_address, result_address, sync_address, props):
|
||||
"""
|
||||
config - must have an entry for 'transforms':array of dicts, which are convertedto configs.
|
||||
"""
|
||||
Transform.__init__(self, feed_address, result_address, sync_address, props)
|
||||
self.transform_address = "tcp://127.0.0.1:{port}".format(port=10104)
|
||||
self.transform_socket = None
|
||||
self.create_transforms(self.config.transforms)
|
||||
|
||||
|
||||
def create_transforms(self, configs):
|
||||
self.transforms = {}
|
||||
for props in configs:
|
||||
class_name = props['class']
|
||||
if(class_name == 'MovingAverage'):
|
||||
mavg = MovingAverage(self.feed_address, self.transform_address, self.sync_address, props)
|
||||
self.transforms[mavg.name] = mavg
|
||||
|
||||
for name, transform in self.transforms.iteritems():
|
||||
self.logger.info("starting {name}".format(name=name))
|
||||
proc = multiprocessing.Process(target=transform.run)
|
||||
proc.start()
|
||||
|
||||
def get_socket(self):
|
||||
|
||||
if(self.transform_socket == None):
|
||||
#create the feed PULL.
|
||||
self.transform_socket = self.context.socket(zmq.PULL)
|
||||
self.transform_socket.connect(self.transform_address)
|
||||
|
||||
def update(self, event):
|
||||
|
||||
state = {}
|
||||
state['feed'] = event
|
||||
|
||||
count = 0
|
||||
while count < len(transforms):
|
||||
message = get_socket().recv
|
||||
data = json.loads(message)
|
||||
state[data['name']] = data
|
||||
|
||||
return state
|
||||
|
||||
|
||||
|
||||
class MovingAverage(Transform):
|
||||
|
||||
def __init__(self, feed_address, result_address, sync_address, props):
|
||||
Transform.__init__(self, feed_address, result_address, sync_address, props)
|
||||
self.events = []
|
||||
|
||||
self.window = datetime.timedelta(days = self.config.get_integer('days'),
|
||||
seconds = self.config.get_integer('seconds'),
|
||||
microseconds = self.config.get_integer('microseconds'),
|
||||
milliseconds = self.config.get_integer('milliseconds'),
|
||||
minutes = self.config.get_integer('minutes'),
|
||||
hours = self.config.get_integer('hours'),
|
||||
weeks = self.config.get_integer('weeks'))
|
||||
|
||||
|
||||
|
||||
|
||||
def update(self, event):
|
||||
self.events.append(event)
|
||||
|
||||
#filter the event list to the window length.
|
||||
self.events = [x for x in self.events if (x.dt - curTick.dt) <= self.window]
|
||||
|
||||
if(len(self.events) == 0):
|
||||
return 0.0
|
||||
|
||||
total = 0.0
|
||||
for event in self.events:
|
||||
total += event['price']
|
||||
|
||||
self.average = total/len(self.events)
|
||||
|
||||
self.state['avg'] = self.average
|
||||
Reference in New Issue
Block a user