mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-11 12:00:50 +08:00
split merging code into a new utility class
This commit is contained in:
+26
-54
@@ -1,5 +1,6 @@
|
||||
|
||||
from data.sources.equity import *
|
||||
from backtest.util import *
|
||||
import time
|
||||
import logging
|
||||
|
||||
@@ -7,29 +8,27 @@ class DataFeed(object):
|
||||
|
||||
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
|
||||
|
||||
self.received_count = 0
|
||||
self.sent_count = 0
|
||||
|
||||
def start_data_workers(self):
|
||||
"""Start a sub-process for each datasource."""
|
||||
|
||||
self.db = db
|
||||
self.data_workers = {}
|
||||
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.data_workers[2] = emt1
|
||||
self.data_buffer[2] = []
|
||||
emt2.start()
|
||||
self.data_workers[2] = emt2
|
||||
|
||||
self.data_buffer = ParallelBuffer(self.data_workers.keys())
|
||||
self.sync_count = subscriber_count + len(self.data_workers)
|
||||
|
||||
|
||||
def start_data_workers(self):
|
||||
"""Start a sub-process for each datasource."""
|
||||
for source_id, source in self.data_workers.iteritems():
|
||||
self.logger.info("starting {id}".format(id=source_id))
|
||||
source.start()
|
||||
self.logger.info("ds processes launched")
|
||||
|
||||
def sync_clients(self):
|
||||
@@ -38,10 +37,10 @@ class DataFeed(object):
|
||||
self.syncservice = self.context.socket(zmq.REP)
|
||||
self.syncservice.bind(self.sync_address)
|
||||
|
||||
|
||||
subscribers = 1
|
||||
total = self.subscriber_count + len(self.data_workers)
|
||||
while subscribers <= total:
|
||||
self.logger.info("sync'ing {count} of {total}".format(count=subscribers, total=total))
|
||||
while subscribers <= self.sync_count:
|
||||
self.logger.info("sync'ing {count} of {total}".format(count=subscribers, total=self.sync_count))
|
||||
# wait for synchronization request
|
||||
msg = self.syncservice.recv()
|
||||
# send synchronization reply
|
||||
@@ -66,6 +65,8 @@ class DataFeed(object):
|
||||
self.feed_socket = self.context.socket(zmq.PUSH)
|
||||
self.feed_socket.bind(self.feed_address)
|
||||
|
||||
self.data_buffer.out_socket = self.feed_socket
|
||||
|
||||
#start the data source workers
|
||||
self.start_data_workers()
|
||||
|
||||
@@ -82,54 +83,25 @@ class DataFeed(object):
|
||||
if(len(self.data_workers) == ds_finished_counter):
|
||||
break
|
||||
else:
|
||||
self.data_buffer[event[u's']].append(event)
|
||||
self.received_count = self.received_count + 1
|
||||
self.send_next()
|
||||
self.data_buffer.append(event[u's'], event)
|
||||
self.data_buffer.send_next()
|
||||
|
||||
|
||||
#drain any remaining messages in the buffer
|
||||
while(self.pending_messages() > 0):
|
||||
self.send_next(drain=True)
|
||||
self.data_buffer.drain()
|
||||
|
||||
#send the DONE message
|
||||
self.feed_socket.send("DONE")
|
||||
|
||||
self.logger.info("received {n} messages, sent {m} messages".format(n=self.received_count, m=self.sent_count))
|
||||
self.logger.info("received {n} messages, sent {m} messages".format(n=self.data_buffer.received_count, m=self.data_buffer.sent_count))
|
||||
self.data_socket.close()
|
||||
self.feed_socket.close()
|
||||
self.context.term()
|
||||
|
||||
|
||||
def send_next(self, drain=False):
|
||||
if(not(self.buffers_full() or drain)):
|
||||
return
|
||||
|
||||
cur = None
|
||||
earliest = None
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if len(events) == 0:
|
||||
continue
|
||||
cur = events
|
||||
if(earliest == None) or (cur[0]['dt'] <= earliest[0]['dt']):
|
||||
earliest = cur
|
||||
|
||||
if(earliest != None):
|
||||
event = earliest.pop(0)
|
||||
self.feed_socket.send(json.dumps(event))
|
||||
self.sent_count += 1
|
||||
|
||||
|
||||
|
||||
def buffers_full(self):
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if (len(events) == 0):
|
||||
return False
|
||||
return True
|
||||
|
||||
def pending_messages(self):
|
||||
total = 0
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
total += len(events)
|
||||
return total
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+78
-43
@@ -4,7 +4,7 @@ import datetime
|
||||
import json
|
||||
import config
|
||||
import multiprocessing
|
||||
from backtest import util
|
||||
from backtest.util import *
|
||||
|
||||
class Transform(object):
|
||||
"""Parent class for feed transforms. Subclass to create a new derived value from the combined feed."""
|
||||
@@ -52,6 +52,9 @@ class Transform(object):
|
||||
sync_socket.close()
|
||||
|
||||
self.logger.info("starting {name} event loop".format(name = self.name))
|
||||
self.run_loop()
|
||||
|
||||
def run_loop(self):
|
||||
|
||||
while True:
|
||||
message = self.feed_socket.recv()
|
||||
@@ -73,8 +76,46 @@ class Transform(object):
|
||||
|
||||
def update(self, event):
|
||||
return {}
|
||||
|
||||
|
||||
class Merge(Transform):
|
||||
class MovingAverage(Transform):
|
||||
|
||||
def __init__(self, feed_address, result_address, sync_address, props, server=False):
|
||||
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 (parse_date(x['dt']) - parse_date(event['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
|
||||
|
||||
return self.state
|
||||
|
||||
|
||||
class MergedTransformsFeed(Transform):
|
||||
""" Merge data feed and array of transform feeds into a single result vector.
|
||||
PULL from feed
|
||||
PULL from child transforms
|
||||
@@ -100,17 +141,46 @@ class Merge(Transform):
|
||||
mavg = MovingAverage(self.feed_address, self.transform_address, self.sync_address, props)
|
||||
self.transforms[mavg.name] = mavg
|
||||
|
||||
self.data_buffer = ParallelBuffer(self.transforms.keys())
|
||||
|
||||
for name, transform in self.transforms.iteritems():
|
||||
self.logger.info("starting {name}".format(name=name))
|
||||
proc = multiprocessing.Process(target=transform.run)
|
||||
proc.start()
|
||||
|
||||
self.buffers = {}
|
||||
for name, transform in self.transforms:
|
||||
self.buffers[name] = []
|
||||
|
||||
def get_socket(self):
|
||||
|
||||
if(self.transform_socket == None):
|
||||
#create the feed PULL.
|
||||
self.transform_socket = self.context.socket(zmq.PULL)
|
||||
self.transform_socket.bind(self.transform_address)
|
||||
return self.transform_socket
|
||||
|
||||
def run_loop(self):
|
||||
|
||||
while True:
|
||||
#get original feed message
|
||||
message = self.feed_socket.recv()
|
||||
self.received_count += 1
|
||||
if(message == "DONE"):
|
||||
self.result_socket.send("DONE")
|
||||
break;
|
||||
event = json.loads(message)
|
||||
self.data_buffer.append(event['name'], event)
|
||||
merged_event = self.data_buffer.merge_next()
|
||||
if(merged_event != None):
|
||||
self.result_socket.send(json.dumps(merged_event))
|
||||
self.sent_count += 1
|
||||
|
||||
self.logger.info("Transform {name} recieved {r} and sent {s}".format(name=self.name, r=self.received_count, s=self.sent_count))
|
||||
|
||||
self.feed_socket.close()
|
||||
self.result_socket.close()
|
||||
self.context.term()
|
||||
|
||||
def update(self, event):
|
||||
|
||||
@@ -118,47 +188,12 @@ class Merge(Transform):
|
||||
state['feed'] = event
|
||||
|
||||
count = 0
|
||||
while count < len(transforms):
|
||||
message = get_socket().recv
|
||||
while count < len(self.transforms):
|
||||
message = self.get_socket().recv()
|
||||
if(message == "DONE"):
|
||||
return "DONE"
|
||||
data = json.loads(message)
|
||||
state[data['name']] = data
|
||||
|
||||
return state
|
||||
|
||||
|
||||
|
||||
class MovingAverage(Transform):
|
||||
|
||||
def __init__(self, feed_address, result_address, sync_address, props, server=False):
|
||||
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 (util.parse_date(x['dt']) - util.parse_date(event['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
|
||||
|
||||
return self.state
|
||||
return state
|
||||
|
||||
Reference in New Issue
Block a user