diff --git a/zipline/gens/composites.py b/zipline/gens/composites.py index 5ee45de6..89eb180f 100644 --- a/zipline/gens/composites.py +++ b/zipline/gens/composites.py @@ -3,36 +3,36 @@ from itertools import tee from zipline.gens.tradegens import SpecificEquityTrades from zipline.gens.utils import roundrobin, hash_args -from zipline.gens.feed import FeedGen -from zipline.gens.merge import MergeGen -from zipline.gens.transform import StatefulTransformGen +from zipline.gens.sort import date_sort +from zipline.gens.merge import merge +from zipline.gens.transform import stateful_transform -def PreTransformLayer(sources, source_args, source_kwargs): +def date_sorted_sources(sources, source_args, source_kwargs): """ Takes a list of generator functions, a list of tuples of positional arguments, and a list of dictionaries of keyword arguments. Packages up all arguments - and passes them into a FeedGen. + and passes them into a date_sort. """ assert len(sources) == len(source_args) == len(source_kwargs) # Package up sources and arguments. arg_bundles = zip(sources, source_args, source_kwargs) - - # Calculate namestring hashes to pass to FeedGen. + + # Calculate namestring hashes to pass to date_sort. namestrings = [source.__name__ + hash_args(*args, **kwargs) for source, args, kwargs in arg_bundles] # Pass each source its arguments. initialized = tuple(source(*args, **kwargs) for source, args, kwargs in arg_bundles) - + stream_in = roundrobin(*initialized) - return FeedGen(stream_in, namestrings) + return date_sort(stream_in, namestrings) -def TransformLayer(feed_stream, tnfms, tnfm_args, tnfm_kwargs): - """ - A generator that takes the expected output of a FeedGen, pipes it +def merged_transforms(sorted_stream, tnfms, tnfm_args, tnfm_kwargs): + """ + A generator that takes the expected output of a date_sort, pipes it through a given set of transforms, and runs the results throught a - MergeGen to output a unified stream. tnfms should be a list of + merge to output a unified stream. tnfms should be a list of pointers to generator functions. tnfm_args should be a list of tuples, representing the arguments to be passed to each transform. tnfm_kwargs should be a list of dictionaries representing keyword @@ -43,38 +43,24 @@ def TransformLayer(feed_stream, tnfms, tnfm_args, tnfm_kwargs): assert len(tnfms) == len(tnfm_args) == len(tnfm_kwargs) # Create a copy of the stream for each transform. - split = tee(feed_stream, len(tnfms)) - - # Package each stream copy with a transform and set of args. Use a list + split = tee(sorted_stream, len(tnfms)) + + # Package each transform with a stream copy and set of args. Use a list # so that we can re-use this for calculating hashes. bundles = zip(split, tnfms, tnfm_args, tnfm_kwargs) - - tnfm_gens = [StatefulTransformGen(stream, tnfm, *args, **kwargs) + + # list comprehension to create transform generators from + # bundles + tnfm_gens = [stateful_transform(stream, tnfm, *args, **kwargs) for stream, tnfm, args, kwargs in bundles] - + # Generate expected hashes for each transform - hashes = [tnfm.__name__ + hash_args(*args, **kwargs) + hashes = [tnfm.__name__ + hash_args(*args, **kwargs) for _, tnfm, args, kwargs in bundles] - + # Roundrobin the outputs of our transforms to create a single flat stream. to_merge = roundrobin(*tnfm_gens) - - # Pipe the stream into MergeGen. - merged = MergeGen(to_merge, hashes) - return merged -if __name__ == "__main__": - from zipline.gens.transform import MovingAverage, Passthrough - - import nose.tools; nose.tools.set_trace() - source = SpecificEquityTrades - feed_out = PreTransformLayer((source,), ((),), ({},)) - - transforms = [MovingAverage, Passthrough] - args = [(datetime.timedelta(days = 1), ['price']), ()] - kwargs = [{}, {}] - - tlayer = TransformLayer(feed_out, transforms, args, kwargs) - - - + # Pipe the stream into merge. + merged = merge(to_merge, hashes) + return merged_transforms diff --git a/zipline/gens/feed.py b/zipline/gens/feed.py deleted file mode 100644 index f029d0ee..00000000 --- a/zipline/gens/feed.py +++ /dev/null @@ -1,129 +0,0 @@ -""" -Generator version of Feed. -""" - -import pytz -import logbook -import pymongo -import types - -from pymongo import ASCENDING -from datetime import datetime, timedelta -from collections import deque, defaultdict - -from zipline import ndict -from zipline.gens.utils import hash_args, assert_datasource_protocol, \ - assert_trade_protocol, assert_datasource_unframe_protocol, \ - assert_feed_protocol - -import zipline.protocol as zp - -def FeedGen(stream_in, source_ids): - """ - A generator that takes a generator and a list of source_ids. We - maintain an internal queue for each id in source_ids. While we - have messages pending from all sources, we pull the earliest - message and yield it. - """ - - assert isinstance(source_ids, list) - - # Set up an internal queue for each expected source. - sources = {} - for id in source_ids: - assert isinstance(id, basestring), "Bad source_id %s" % source_id - sources[id] = deque() - - # Process incoming streams. - for message in stream_in: - # Incoming messages should be the output of DATASOURCE_UNFRAME. - assert_datasource_unframe_protocol(message), \ - "Bad message in FeedGen: %s" % message - - # Only allow messages from sources we expect. - assert message.source_id in sources, "Unexpected source: %s" % message - - sources[message.source_id].append(message) - - # Only pop messages when we have a pending message from - # all datasources. Stop if all sources have signalled done. - - while full(sources) and not done(sources): - message = pop_oldest(sources) - assert_feed_protocol(message) - yield message - - # We should have only a done message left in each queue. - for queue in sources.itervalues(): - assert len(queue) == 1, "Bad queue in FeedGen on exit: %s" % queue - assert queue[0].dt == "DONE", \ - "Bad last message in FeedGen on exit: %s" % queue - -def full(sources): - """ - Feed is full when every internal queue has at least one - message. Note that this include DONE messages, so done(sources) is - True only if full(sources). - """ - assert isinstance(sources, dict) - return all( (queue_is_full(source) for source in sources.itervalues()) ) - -def queue_is_full(queue): - assert isinstance(queue, deque) - return len(queue) > 0 - -def done(sources): - """Feed is done when all internal queues have only a "DONE" message.""" - assert isinstance(sources, dict) - return all( (queue_is_done(source) for source in sources.itervalues()) ) - -def queue_is_done(queue): - assert isinstance(queue, deque) - if len(queue) == 0: - return False - if queue[0].dt == "DONE": - assert len(queue) == 1, "Message after DONE in FeedGen: %s" % queue - return True - else: - return False - -def pop_oldest(sources): - - oldest_event = None - - # Iterate over the dict, checking internal queues for the oldest - # pending event. - - for queue in sources.itervalues(): - current_event = queue[0] - # Skip queues that are done. - if current_event.dt == "DONE": - continue - # Any event is older than nothing. - elif oldest_event == None: - oldest_event = current_event - # Keep the older event. Break ties by source_id. This will - # trip an assert if we have duplicate sources. - else: - oldest_event = older(oldest_event, current_event) - - # Pop the oldest event we found from its queue and return it. - return sources[oldest_event.source_id].popleft() - -# Return the event with the older timestamp. Break ties by source_id. -def older(oldest, current): - assert isinstance(oldest, ndict) - assert isinstance(oldest, ndict) - - # Try to compare by dt. - if oldest.dt < current.dt: - return oldest - elif oldest.dt > current.dt: - return current - # Break ties by source_id. - elif oldest.source_id < current.source_id: - return oldest - elif oldest.source_id > current.source_id: - return current - else: - assert False, "Duplicate event" diff --git a/zipline/gens/merge.py b/zipline/gens/merge.py index 52adcb84..4778ed5b 100644 --- a/zipline/gens/merge.py +++ b/zipline/gens/merge.py @@ -2,32 +2,24 @@ Generator version of Merge. """ -import pytz -import logbook -import pymongo -import types - -from pymongo import ASCENDING -from datetime import datetime, timedelta -from collections import deque, defaultdict +from collections import deque from zipline import ndict -from zipline.gens.utils import hash_args, assert_datasource_protocol, \ - assert_trade_protocol, assert_datasource_unframe_protocol, assert_merge_protocol +from zipline.gens.utils import hash_args, \ + assert_merge_protocol -import zipline.protocol as zp -def MergeGen(stream_in, tnfm_ids): +def merge(stream_in, tnfm_ids): """ A generator that takes a generator and a list of source_ids. We - maintain an internal queue for each id in source_ids. Once we + maintain an internal queue for each id in source_ids. Once we have a message from every queue, we pop an event from each queue - and merge them together into an event. We raise an error if we + and merge them together into an event. We raise an error if we do not receive the same number of events from all sources. """ assert isinstance(tnfm_ids, list) - + # Set up an internal queue for each expected source. tnfms = {} for id in tnfm_ids: @@ -37,28 +29,28 @@ def MergeGen(stream_in, tnfm_ids): # Process incoming streams. for message in stream_in: assert isinstance(message, tuple), \ - "Bad message in MergeGen: %s" %message + "Bad message in merge: %s" %message assert len(message) == 2 id, value = message assert id in tnfm_ids, \ "Message from unexpected tnfm: %s, %s" % (id, tnfm_ids) - assert isinstance(value, ndict), "Bad message in MergeGen: %s" %message + assert isinstance(value, ndict), "Bad message in merge: %s" %message tnfms[id].append(value) # Only pop messages when we have a pending message from # all datasources. Stop if all sources have signalled done. - while full(tnfms) and not done(tnfms): + while ready(tnfms) and not done(tnfms): message = merge_one(tnfms) assert_merge_protocol(tnfm_ids, message) yield message - # We should have only a done message left in each queue. + # We should have only a done message left in each queue. for queue in tnfms.itervalues(): - assert len(queue) == 1, "Bad queue in MergeGen on exit: %s" % queue + assert len(queue) == 1, "Bad queue in merge on exit: %s" % queue assert queue[0].dt == "DONE", \ - "Bad last message in MergeGen on exit: %s" % queue + "Bad last message in merge on exit: %s" % queue def merge_one(sources): output = ndict() @@ -68,22 +60,22 @@ def merge_one(sources): return output -#TODO: This is replicated in feed. Probably should be one source file. -def full(sources): +#TODO: This is replicated in sort. Probably should be one source file. +def ready(sources): """ - Feed is full when every internal queue has at least one message. Note that - this include DONE messages, so done(sources) is True only if full(sources). + Feed is ready when every internal queue has at least one message. Note that + this include DONE messages, so done(sources) is True only if ready(sources). """ assert isinstance(sources, dict) - return all( (queue_is_full(source) for source in sources.itervalues()) ) + return all( (queue_is_ready(source) for source in sources.itervalues()) ) -def queue_is_full(queue): - assert isinstance(queue, deque) +def queue_is_ready(queue): + assert isinstance(queue, deque) return len(queue) > 0 def done(sources): """Feed is done when all internal queues have only a "DONE" message.""" - assert isinstance(sources, dict) + assert isinstance(sources, dict) return all( (queue_is_done(source) for source in sources.itervalues()) ) def queue_is_done(queue): @@ -91,11 +83,7 @@ def queue_is_done(queue): if len(queue) == 0: return False if queue[0].dt == "DONE": - assert len(queue) == 1, "Message after DONE in FeedGen: %s" % queue + assert len(queue) == 1, "Message after DONE in date_sort: %s" % queue return True else: return False - - - - diff --git a/zipline/gens/mongods.py b/zipline/gens/mongods.py index adcb0ed3..aa22d4c1 100644 --- a/zipline/gens/mongods.py +++ b/zipline/gens/mongods.py @@ -7,10 +7,10 @@ import logbook import pymongo from pymongo import ASCENDING -from datetime import datetime, timedelta +from datetime import datetime from zipline import ndict -from zipline.gens.utils import hash_args, assert_datasource_protocol, \ +from zipline.gens.utils import hash_args, \ assert_trade_protocol import zipline.protocol as zp @@ -22,12 +22,12 @@ def MongoTradeHistoryGen(collection, filter, start_date, end_date): start, and end. The output is also packaged with a unique source_id string for downstream sorting """ - + assert isinstance(collection, pymongo.collection.Collection) assert isinstance(filter, dict) assert isinstance(start_date, (datetime)) assert isinstance(end_date, (datetime)) - + # Set up internal iterator. This outputs raw dictionaries. iterator = create_pymongo_iterator(collection, filter, start_date, end_date) @@ -46,13 +46,13 @@ def MongoTradeHistoryGen(collection, filter, start_date, end_date): payload = ndict(event) assert_trade_protocol(payload) yield payload - + def create_pymongo_iterator(collection, filter, start_date, end_date): """ Returns an iterator that spits out raw objects loaded from a - MongoDB collection. + MongoDB collection. - See the comments on :py:class:`zipline.messaging.DataSource` + See the comments on :py:class:`zipline.messaging.DataSource` for expected content of filter. """ log = logbook.Logger("MongoDBQuery") @@ -68,7 +68,7 @@ def create_pymongo_iterator(collection, filter, start_date, end_date): assert isinstance(value, list) sid_range = {'sid':{'$in':value}} spec.update(sid_range) - + # limit the data to the date range [start, end], inclusive date_range = {'dt':{'$gte': start_date, '$lte': end_date}} spec.update(date_range) @@ -78,7 +78,7 @@ def create_pymongo_iterator(collection, filter, start_date, end_date): # In our collection, load all objects matching spec. Of those # objects, get only the fields matching fields, and return the # loaded objects sorted by dt from least to greatest. - + cursor = collection.find( fields = fields, spec = spec, @@ -92,11 +92,5 @@ def create_pymongo_iterator(collection, filter, start_date, end_date): # Set up the iterator iterator = iter(cursor) log.info("MongoDataSource iterator ready") - + return iterator - - - - - - diff --git a/zipline/gens/scratch.py b/zipline/gens/scratch.py deleted file mode 100644 index 2de87e44..00000000 --- a/zipline/gens/scratch.py +++ /dev/null @@ -1,36 +0,0 @@ - -# Inside Client -def __init__(self, addresses, ...): - self.pull_socket = ... - - self.control_socket ... - -def run(self): - for message in gen_from_pull(self.pull_socket): - #Do things with messages. - heartbeat() - - signal_done() - sys.exit(0) - -# Inside Merge -def __init__(self, addresses, source_ids ...): - self.poller = ... # Poller on multiple xforms, single socket. - - self.processor = ... # Generator that - - self.push_socket = ... # Outbound socket - -def run(self): - - incoming = gen_from_poll(self.poller)# Receives messages from all xforms. - - processed = self.processor(incoming, source_ids) # Maintains internal queues and merges. - - for message in self.processed: - heartbeat() - self.push_socket.send(message) - -# Inside - - diff --git a/zipline/gens/test_feed.py b/zipline/gens/test_feed.py index b5007b91..bfff6742 100644 --- a/zipline/gens/test_feed.py +++ b/zipline/gens/test_feed.py @@ -11,16 +11,16 @@ from datetime import datetime, timedelta from collections import deque from zipline import ndict -from zipline.gens.feed import FeedGen, full, done, queue_is_full,queue_is_done,\ +from zipline.gens.sort import date_sort, ready, done, queue_is_ready,queue_is_done,\ pop_oldest from zipline.gens.utils import hash_args, assert_datasource_protocol,\ assert_trade_protocol, alternate from zipline.gens.tradegens import date_gen, SpecificEquityTrades -from zipline.gens.composites import PreTransformLayer +from zipline.gens.composites import date_sorted_sources import zipline.protocol as zp -class FeedHelperTestCase(TestCase): +class HelperTestCase(TestCase): def setUp(self): pass @@ -30,23 +30,23 @@ class FeedHelperTestCase(TestCase): def test_individual_queue_logic(self): queue = deque() - # Empty queues are neither done nor full. - assert not queue_is_full(queue) + # Empty queues are neither done nor ready. + assert not queue_is_ready(queue) assert not queue_is_done(queue) queue.append(to_dt('foo')) - assert queue_is_full(queue) + assert queue_is_ready(queue) assert not queue_is_done(queue) queue.appendleft(to_dt('DONE')) - assert queue_is_full(queue) + assert queue_is_ready(queue) # Checking done when we have a message after done will trip an assert. self.assertRaises(AssertionError, queue_is_done, queue) queue.pop() - assert queue_is_full(queue) + assert queue_is_ready(queue) assert queue_is_done(queue) def test_pop_logic(self): @@ -55,35 +55,35 @@ class FeedHelperTestCase(TestCase): for id in ids: sources[id] = deque() - assert not full(sources) + assert not ready(sources) assert not done(sources) - # All sources must have a message to be full/done + # All sources must have a message to be ready/done sources['a'].append(to_dt("datetime")) - assert not full(sources) + assert not ready(sources) assert not done(sources) sources['a'].pop() for id in ids: sources[id].append(to_dt("datetime")) - assert full(sources) + assert ready(sources) assert not done(sources) for id in ids: sources[id].appendleft(to_dt("DONE")) # ["DONE", message] will trip an assert in queue_is_done. - assert full(sources) + assert ready(sources) self.assertRaises(AssertionError, done, sources) for id in ids: sources[id].pop() - assert full(sources) + assert ready(sources) assert done(sources) -class FeedGenTestCase(TestCase): +class DateSortTestCase(TestCase): def setUp(self): pass @@ -91,13 +91,13 @@ class FeedGenTestCase(TestCase): def tearDown(self): pass - def run_FeedGen(self, events, expected, source_ids): + def run_date_sort(self, events, expected, source_ids): """ Take a list of events, their source_ids, and an expected sorting. - Assert that FeedGen's output agrees with expected. + Assert that date_sort's output agrees with expected. """ - feed_gen = FeedGen(events, source_ids) - l = list(feed_gen) + sort_gen = date_sort(events, source_ids) + l = list(sort_gen) assert l == expected def test_single_source(self): @@ -118,7 +118,7 @@ class FeedGenTestCase(TestCase): event_gen = (e for e in events) - self.run_FeedGen(event_gen, expected, source_ids) + self.run_date_sort(event_gen, expected, source_ids) def test_multi_source(self): source_ids = ['a', 'b'] @@ -147,14 +147,14 @@ class FeedGenTestCase(TestCase): # Alternating between a and b. interleaved = alternate(iter(events_a), iter(events_b)) - self.run_FeedGen(interleaved, expected, source_ids) + self.run_date_sort(interleaved, expected, source_ids) # All of a, then all of b. sequential = chain(iter(events_a), iter(events_b)) - self.run_FeedGen(sequential, expected, source_ids) + self.run_date_sort(sequential, expected, source_ids) - def test_full_feed_layer(self): + def test_sorted_sources(self): filter = [1,2] #Set up source a. One hour between events. @@ -196,12 +196,12 @@ class FeedGenTestCase(TestCase): expected_ids = ["SpecificEquityTrades" + hash_args(*args, **kwargs) for args, kwargs in zip_args] - # Pipe our sources into feed. - feed_out = PreTransformLayer(sources, source_args, source_kwargs) + # Pipe our sources into sort. + sort_out = date_sorted_sources(sources, source_args, source_kwargs) - # Read all the values from feed and assert that they arrive in + # Read all the values from sort and assert that they arrive in # the correct sorting with the expected hash values. - to_list = list(feed_out) + to_list = list(sort_out) copy = to_list[:] for e in to_list: # All events should match one of our expected source_ids. diff --git a/zipline/gens/tradegens.py b/zipline/gens/tradegens.py index fb0b3f48..c7ee74f8 100644 --- a/zipline/gens/tradegens.py +++ b/zipline/gens/tradegens.py @@ -1,12 +1,16 @@ +""" +Tools to generate trade events without a backing store. Useful for testing +and zipline development +""" import random -from itertools import chain, repeat, cycle, ifilter, izip +from itertools import chain, cycle, ifilter, izip from datetime import datetime, timedelta from zipline.utils.factory import create_trade from zipline.gens.utils import hash_args, mock_done def date_gen(start = datetime(2012, 6, 6, 0), - delta = timedelta(minutes = 1), + delta = timedelta(minutes = 1), count = 100): """ Utility to generate a stream of dates. @@ -29,13 +33,13 @@ def mock_volumes(count, rand = False): """ Utility to generate a set of volumes. By default cycles through values from 100 to 1000, incrementing by 50. Optional - flag to give random values between 100 and 1000. + flag to give random values between 100 and 1000. """ if rand: return (random.randrange(100, 1000) for i in xrange(count)) else: return ((i * 50)%900 + 100 for i in xrange(count)) - + def fuzzy_dates(count = 500): """ Add +-10 seconds to each event from a date_gen. Note that this @@ -43,7 +47,7 @@ def fuzzy_dates(count = 500): separation of events. """ for date in date_gen(count = count): - yield date + timedelta(seconds = random.randint(-10, 10)) + yield date + timedelta(seconds = random.randint(-10, 10)) def SpecificEquityTrades(*args, **config): """ @@ -61,7 +65,7 @@ def SpecificEquityTrades(*args, **config): delta = config.get('delta', timedelta(minutes = 1)) # Default to None for event_list and filter. - event_list = config.get('event_list') + event_list = config.get('event_list') filter = config.get('filter') arg_string = hash_args(*args, **config) @@ -101,7 +105,7 @@ def SpecificEquityTrades(*args, **config): def RandomEquityTrades(*args, **config): # We shouldn't get any positional args. assert args == () - + count = config.get('count', 500) sids = config.get('sids', [1,2]) filter = config.get('filter') @@ -112,9 +116,9 @@ def RandomEquityTrades(*args, **config): sids = cycle(sids) arg_gen = izip(sids, prices, volumes, dates) - + unfiltered = (create_trade(*args) for args in arg_gen) - + if filter: filtered = ifilter(lambda event: event.sid in filter, unfiltered) else: diff --git a/zipline/gens/transform.py b/zipline/gens/transform.py index 137b9d69..a03a841a 100644 --- a/zipline/gens/transform.py +++ b/zipline/gens/transform.py @@ -1,64 +1,56 @@ """ Generator versions of transforms. """ -import random -import pytz -import logbook -import pymongo import types -from pymongo import ASCENDING -from datetime import datetime, timedelta +from datetime import datetime from collections import deque, defaultdict from numbers import Number -from itertools import izip from zipline import ndict -from zipline.gens.tradegens import date_gen -from zipline.gens.utils import assert_feed_unframe_protocol, \ +from zipline.gens.utils import assert_sort_unframe_protocol, \ assert_transform_protocol, hash_args -import zipline.protocol as zp - class Passthrough(object): """ - Trivial function for forwarding events. + Trivial class for forwarding events. """ def __init__(self): pass - def update(self, event): + def update(self, event): assert isinstance(event, ndict),"Bad event in Passthrough: %s" % event assert event.has_key('sid'), "No sid in Passthrough: %s" % event assert event.has_key('dt'), "No dt in Passthorughz: %s" % event return event -def FunctionalTransformGen(stream_in, fun, *args, **kwargs): +def functional_transform(stream_in, func, *args, **kwargs): """ Generic transform generator that takes each message from an in-stream and yields the output of a function on that message. Not sure how useful this will be in reality, but good for testing. """ - - # TODO: Distinguish between functions and classes in hash_args. - # As implemented we will get assertion errors if a function and - # stateful class have the same name, which may or may not be - # what we want. + assert isinstance(func, types.FunctionType), \ + "Functional" + namestring = func.__name__ + hash_args(*args, **kwargs) - namestring = fun.__name__ + hash_args(*args, **kwargs) - for message in stream_in: - assert_feed_unframe_protocol(message) - out_value = fun(message, *args, **kwargs) + assert_sort_unframe_protocol(message) + out_value = func(message, *args, **kwargs) assert_transform_protocol(out_value) yield(namestring, out_value) - -def StatefulTransformGen(stream_in, tnfm_class, *args, **kwargs): + +def stateful_transform(stream_in, tnfm_class, *args, **kwargs): """ Generic transform generator that takes each message from an in-stream - and feeds it to a state class. For each call to update, the state + and sorts it to a state class. For each call to update, the state class must produce a message to be fed downstream. - """ + """ + + assert isinstance(tnfm_class, (types.ObjectType, types.ClassType)), \ + "Stateful transform requires a class." + assert tnfm_class.__dict__.has_key('update'), \ + "Stateful transform requires the class to have an update method" # Create an instance of our transform class. state = tnfm_class(*args, **kwargs) @@ -67,7 +59,7 @@ def StatefulTransformGen(stream_in, tnfm_class, *args, **kwargs): namestring = tnfm_class.__name__ + hash_args(*args, **kwargs) for message in stream_in: - assert_feed_unframe_protocol(message) + assert_sort_unframe_protocol(message) out_value = state.update(message) assert_transform_protocol(out_value) yield (namestring, out_value) @@ -78,7 +70,7 @@ class MovingAverage(object): Upon receipt of each message we update the corresponding window and return the calculated average. """ - + def __init__(self, delta, fields): self.delta = delta self.fields = fields @@ -86,38 +78,38 @@ class MovingAverage(object): # No way to pass arguments to the defaultdict factory, so we # need to define a method to generate the correct EventWindows. self.sid_windows = defaultdict(self.create_window) - + def create_window(self): """Factory method for self.sid_windows.""" return EventWindow(self.delta, self.fields) - + def update(self, event): """ Update the event window for this event's sid. Return an ndict from tracked fields to averages. """ - + assert isinstance(event, ndict),"Bad event in MovingAverage: %s" % event assert event.has_key('sid'), "No sid in MovingAverage: %s" % event assert event.has_key('dt'), "No dt in MovingAverage: %s" % event - + output = ndict({'sid': event.sid, 'dt': event.dt}) # This will create a new EventWindow if this is the first # message for this sid. window = self.sid_windows[event.sid] window.update(event) averages = window.get_averages() - + # Return the calculated averages along with output.merge(averages) return output - + class EventWindow(object): """ Maintains a list of events that are within a certain timedelta of the most recent tick. The expected use of this class is to track events associated with a single sid. We provide simple - functionality for averages, but anything more complicated + functionality for averages, but anything more complicated should be handled by a containing class. """ @@ -126,17 +118,17 @@ class EventWindow(object): self.delta = delta self.fields = fields self.totals = defaultdict(float) - + def __len__(self): return len(self.ticks) - + def update(self, event): self.assert_well_formed(event) # Add new event and increment totals. self.ticks.append(event) for field in self.fields: self.totals[field] += event[field] - + # We return a list of all out-of-range events we removed. out_of_range = [] @@ -144,7 +136,7 @@ class EventWindow(object): # newest oldest # | | # V V - + while (self.ticks[-1].dt - self.ticks[0].dt) >= self.delta: # popleft removes and returns ticks[0] popped = self.ticks.popleft() @@ -155,7 +147,7 @@ class EventWindow(object): out_of_range.append(popped) return out_of_range - + def average(self, field): assert field in self.fields if len(self.ticks) == 0: @@ -183,30 +175,6 @@ class EventWindow(object): "Events arrived out of order in EventWindow: %s -> %s" % (event, self.ticks[0]) for field in self.fields: assert event.has_key(field), \ - "Event missing [%s] in EventWindow" % field + "Event missing [%s] in EventWindow" % field assert isinstance(event[field], Number), \ "Got %s for %s in EventWindow" % (event[field], field) - -# if __name__ == "__main__": - -# def make_event(**kwargs): -# e = ndict() -# for key, value in kwargs.iteritems(): -# e[key] = value -# return e - -# dates = date_gen(delta = timedelta(hours = 12)) -# events = ( -# make_event( -# sid = 'foo', price = random.random(), -# dt = date, -# type = zp.DATASOURCE_TYPE.TRADE, -# source_id = 'ds', -# vol = i -# ) -# for date, i in izip(dates, xrange(100)) -# ) - -# gen = MovingAverageTransformGen(events, 1, ['price', 'vol']) - - diff --git a/zipline/gens/utils.py b/zipline/gens/utils.py index c10cffdf..e2f859cb 100644 --- a/zipline/gens/utils.py +++ b/zipline/gens/utils.py @@ -2,7 +2,7 @@ import pytz import numbers from hashlib import md5 -from datetime import datetime, timedelta +from datetime import datetime from itertools import izip_longest from zipline import ndict from zipline.protocol import DATASOURCE_TYPE @@ -18,7 +18,7 @@ def mock_raw_event(sid, dt): def mock_done(source_id): return ndict({'dt': "DONE", "source_id" : source_id, 'type' : 0}) - + def alternate(g1, g2): """Specialized version of roundrobin for just 2 generators.""" for e1, e2 in izip_longest(g1, g2): @@ -30,7 +30,7 @@ def alternate(g1, g2): def roundrobin(*args): """ Takes N generators, pulling one element off each until all inputs - are empty. + are empty. """ for elem_tuple in izip_longest(*args): for value in elem_tuple: @@ -43,11 +43,11 @@ def hash_args(*args, **kwargs): arg_string = '_'.join([str(arg) for arg in args]) kwarg_string = '_'.join([str(key) + '=' + str(value) for key, value in kwargs.iteritems()]) combined = ':'.join([arg_string, kwarg_string]) - + hasher = md5() hasher.update(combined) return hasher.hexdigest() - + def assert_datasource_protocol(event): """Assert that an event meets the protocol for datasource outputs.""" @@ -59,7 +59,7 @@ def assert_datasource_protocol(event): if not event.type == DATASOURCE_TYPE.DONE: assert isinstance(event.dt, datetime) assert event.dt.tzinfo == pytz.utc - + def assert_trade_protocol(event): """Assert that an event meets the protocol for datasource TRADE outputs.""" assert_datasource_protocol(event) @@ -77,15 +77,15 @@ def assert_datasource_unframe_protocol(event): assert isinstance(event.source_id, basestring) assert event.type in DATASOURCE_TYPE assert event.has_key('dt') - -def assert_feed_protocol(event): + +def assert_sort_protocol(event): """Assert that an event is valid input to zp.FEED_FRAME.""" assert isinstance(event, ndict) assert isinstance(event.source_id, basestring) assert event.type in DATASOURCE_TYPE assert event.has_key('dt') - -def assert_feed_unframe_protocol(event): + +def assert_sort_unframe_protocol(event): """Same as above.""" assert isinstance(event, ndict) assert isinstance(event.source_id, basestring) @@ -93,11 +93,10 @@ def assert_feed_unframe_protocol(event): assert event.has_key('dt') def assert_transform_protocol(event): - """Transforms should return an ndict to be merged by MergeGen.""" + """Transforms should return an ndict to be merged by merge.""" assert isinstance(event, ndict) def assert_merge_protocol(tnfm_ids, message): """Merge should output an ndict with a field for each id in its transform set.""" assert isinstance(message, ndict) assert set(tnfm_ids) == set(message.keys()) -