diff --git a/Algorithm.Python/PythonDictionaryFeatureRegressionAlgorithm.py b/Algorithm.Python/PythonDictionaryFeatureRegressionAlgorithm.py index 9cf68fe65466..d7f56c21f521 100644 --- a/Algorithm.Python/PythonDictionaryFeatureRegressionAlgorithm.py +++ b/Algorithm.Python/PythonDictionaryFeatureRegressionAlgorithm.py @@ -52,6 +52,15 @@ def test_slice_dictionary(self): if spy is None: raise AssertionError('SPY is not in Slice') + if slice.contains_key(None): + raise AssertionError('Slice.contains_key(None) should return False instead of throwing') + + if slice.get(None) is not None: + raise AssertionError('Slice.get(None) should return None instead of throwing') + + if slice.bars.contains_key(None): + raise AssertionError('TradeBars.contains_key(None) should return False instead of throwing') + for symbol, bar in slice.bars.items(): self.plot(symbol, 'Price', bar.close) @@ -74,6 +83,16 @@ def test_securities_dictionary(self): if aapl is not None: raise AssertionError('aapl is not None') + # A None key should behave like a missing key instead of throwing, + # e.g. when a symbol field is only assigned later in the algorithm + none_symbol = None + price = self.securities[none_symbol].price if self.securities.contains_key(none_symbol) else None + if price is not None: + raise AssertionError('Securities.contains_key(None) should return False instead of throwing') + + if self.securities.get(none_symbol) is not None: + raise AssertionError('Securities.get(None) should return None instead of throwing') + for symbol, security in self.securities.items(): self.plot(symbol, 'Price', security.price) @@ -95,6 +114,12 @@ def test_portfolio_dictionary(self): if aapl is not None: raise AssertionError('aapl is not None') + if self.portfolio.contains_key(None): + raise AssertionError('Portfolio.contains_key(None) should return False instead of throwing') + + if self.portfolio.get(None) is not None: + raise AssertionError('Portfolio.get(None) should return None instead of throwing') + for symbol, holdings in self.portfolio.items(): msg = f'{symbol}: {holdings.leverage}' diff --git a/Common/Data/Market/OptionChains.cs b/Common/Data/Market/OptionChains.cs index 75f67c8315ed..54b289b7b19f 100644 --- a/Common/Data/Market/OptionChains.cs +++ b/Common/Data/Market/OptionChains.cs @@ -113,6 +113,11 @@ public override bool Remove(KeyValuePair item) private static Symbol GetCanonicalOptionSymbol(Symbol symbol) { + if (ReferenceEquals(symbol, null)) + { + return null; + } + if (symbol.SecurityType.HasOptions()) { return Symbol.CreateCanonicalOption(symbol); diff --git a/Common/Data/Slice.cs b/Common/Data/Slice.cs index 18c5a3ade2fb..cf678e84cf90 100644 --- a/Common/Data/Slice.cs +++ b/Common/Data/Slice.cs @@ -527,7 +527,7 @@ public T Get(Symbol symbol) /// True if this instance contains data for the symbol, false otherwise public override bool ContainsKey(Symbol symbol) { - return _data.Value.ContainsKey(symbol); + return !ReferenceEquals(symbol, null) && _data.Value.ContainsKey(symbol); } /// @@ -540,7 +540,7 @@ public override bool TryGetValue(Symbol symbol, out dynamic data) { data = null; SymbolData symbolData; - if (_data.Value.TryGetValue(symbol, out symbolData)) + if (!ReferenceEquals(symbol, null) && _data.Value.TryGetValue(symbol, out symbolData)) { data = symbolData.GetData(); return data != null; diff --git a/Common/Securities/CashBook.cs b/Common/Securities/CashBook.cs index 620950e03bac..d8271424e8f3 100644 --- a/Common/Securities/CashBook.cs +++ b/Common/Securities/CashBook.cs @@ -279,7 +279,7 @@ public bool Remove(KeyValuePair item) /// Key. public override bool ContainsKey(string symbol) { - return _currencies.ContainsKey(symbol); + return !ReferenceEquals(symbol, null) && _currencies.ContainsKey(symbol); } /// @@ -291,6 +291,11 @@ public override bool ContainsKey(string symbol) /// Value. public override bool TryGetValue(string symbol, out Cash value) { + if (ReferenceEquals(symbol, null)) + { + value = null; + return false; + } return _currencies.TryGetValue(symbol, out value); } diff --git a/Common/Securities/Positions/SecurityPositionGroupModel.cs b/Common/Securities/Positions/SecurityPositionGroupModel.cs index 2e26ad1d6744..19f091f38f8b 100644 --- a/Common/Securities/Positions/SecurityPositionGroupModel.cs +++ b/Common/Securities/Positions/SecurityPositionGroupModel.cs @@ -246,6 +246,11 @@ private void ResolvePositionGroups() /// True if a group with the specified key was found, false otherwise public override bool TryGetValue(PositionGroupKey key, out IPositionGroup value) { + if (ReferenceEquals(key, null)) + { + value = null; + return false; + } return Groups.TryGetGroup(key, out value); } } diff --git a/Common/Securities/SecurityManager.cs b/Common/Securities/SecurityManager.cs index eba37583b819..03ad9b133bf7 100644 --- a/Common/Securities/SecurityManager.cs +++ b/Common/Securities/SecurityManager.cs @@ -147,6 +147,10 @@ public bool Contains(KeyValuePair pair) /// Bool true if contains this symbol pair public override bool ContainsKey(Symbol symbol) { + if (ReferenceEquals(symbol, null)) + { + return false; + } lock (_securityManager) { return _completeSecuritiesCollection.ContainsKey(symbol); @@ -252,6 +256,11 @@ public ICollection Keys /// True on successfully locating the security object public override bool TryGetValue(Symbol symbol, out Security security) { + if (ReferenceEquals(symbol, null)) + { + security = null; + return false; + } lock (_securityManager) { return _completeSecuritiesCollection.TryGetValue(symbol, out security); diff --git a/Common/Util/BaseExtendedDictionary.cs b/Common/Util/BaseExtendedDictionary.cs index effef7c89ebd..926637baff51 100644 --- a/Common/Util/BaseExtendedDictionary.cs +++ b/Common/Util/BaseExtendedDictionary.cs @@ -83,6 +83,11 @@ public BaseExtendedDictionary(IEnumerable data, Func keySe /// true if the key was found; otherwise, false public override bool TryGetValue(TKey key, out TValue value) { + if (ReferenceEquals(key, null)) + { + value = default; + return false; + } return Dictionary.TryGetValue(key, out value); } @@ -170,7 +175,7 @@ public virtual void Add(KeyValuePair item) /// true if the dictionary contains an element with the specified key; otherwise, false public override bool ContainsKey(TKey key) { - return Dictionary.ContainsKey(key); + return !ReferenceEquals(key, null) && Dictionary.ContainsKey(key); } /// diff --git a/Tests/Common/ExtendedDictionaryTests.cs b/Tests/Common/ExtendedDictionaryTests.cs index dbd7de14dbb0..bd3d287c4fbc 100644 --- a/Tests/Common/ExtendedDictionaryTests.cs +++ b/Tests/Common/ExtendedDictionaryTests.cs @@ -16,9 +16,13 @@ using NUnit.Framework; using Python.Runtime; using QuantConnect.Statistics; +using System; using System.Collections.Generic; using System.Linq; +using QuantConnect.Data; using QuantConnect.Data.Market; +using QuantConnect.Securities; +using QuantConnect.Securities.Positions; namespace QuantConnect.Tests.Common { @@ -229,6 +233,39 @@ def set(dictionary, key, value): Assert.IsInstanceOf(exception.InnerException); } + private static IEnumerable NullKeyTestCases() + { + var time = new DateTime(2025, 1, 1); + var securities = new SecurityManager(new TimeKeeper(time, TimeZones.NewYork)); + + yield return new TestCaseData(securities).SetArgDisplayNames(nameof(SecurityManager)); + yield return new TestCaseData(new SecurityPortfolioManager(securities, new SecurityTransactionManager(null, securities), new AlgorithmSettings())) + .SetArgDisplayNames(nameof(SecurityPortfolioManager)); + yield return new TestCaseData(new CashBook()).SetArgDisplayNames(nameof(CashBook)); + yield return new TestCaseData(new Slice(time, new List(), time)).SetArgDisplayNames(nameof(Slice)); + yield return new TestCaseData(new DataDictionary()).SetArgDisplayNames("DataDictionary"); + yield return new TestCaseData(new TradeBars()).SetArgDisplayNames(nameof(TradeBars)); + yield return new TestCaseData(new OptionChains()).SetArgDisplayNames(nameof(OptionChains)); + yield return new TestCaseData(new FuturesChains()).SetArgDisplayNames(nameof(FuturesChains)); + yield return new TestCaseData(new UniverseManager()).SetArgDisplayNames(nameof(UniverseManager)); + yield return new TestCaseData(new SecurityPositionGroupModel()).SetArgDisplayNames(nameof(SecurityPositionGroupModel)); + } + + [TestCaseSource(nameof(NullKeyTestCases))] + public void DictionariesHandleNullKeysGracefully(object dictionary) + { + AssertNullKeyIsHandledGracefully((dynamic)dictionary); + } + + private static void AssertNullKeyIsHandledGracefully(ExtendedDictionary dictionary) + where TKey : class + where TValue : class + { + Assert.IsFalse(dictionary.ContainsKey(null)); + Assert.IsFalse(dictionary.TryGetValue(null, out _)); + Assert.IsNull(dictionary.get(null)); + } + private class TestDictionary : ExtendedDictionary { private readonly Dictionary _data = new();