diff --git a/src/ScriptEngine/Machine/Contexts/AttachedScriptsFactory.cs b/src/ScriptEngine/Machine/Contexts/AttachedScriptsFactory.cs index 82f7db917..39e701d6a 100755 --- a/src/ScriptEngine/Machine/Contexts/AttachedScriptsFactory.cs +++ b/src/ScriptEngine/Machine/Contexts/AttachedScriptsFactory.cs @@ -5,6 +5,7 @@ This Source Code Form is subject to the terms of the at http://mozilla.org/MPL/2.0/. ----------------------------------------------------------*/ using System; +using System.Collections.Concurrent; using System.Collections.Generic; using System.Text; using OneScript.Sources; @@ -22,39 +23,27 @@ namespace ScriptEngine.Machine.Contexts { public class AttachedScriptsFactory { - private readonly Dictionary _loadedModules; - private readonly Dictionary _fileHashes; + // Сценарии подключаются и из фоновых заданий: регистрация под блокировкой, + // модули читаются без нее при каждом создании объекта + private readonly ConcurrentDictionary _loadedModules; + private readonly ConcurrentDictionary _fileHashes; + private readonly object _registrationLock = new object(); private readonly ScriptingEngine _engine; internal AttachedScriptsFactory(ScriptingEngine engine) { - _loadedModules = new Dictionary(StringComparer.InvariantCultureIgnoreCase); - _fileHashes = new Dictionary(StringComparer.InvariantCultureIgnoreCase); + _loadedModules = new ConcurrentDictionary(StringComparer.InvariantCultureIgnoreCase); + _fileHashes = new ConcurrentDictionary(StringComparer.InvariantCultureIgnoreCase); _engine = engine; } private ITypeManager TypeManager => _engine.TypeManager; - static string GetMd5Hash(MD5 md5Hash, string input) + // Хеш нужен только чтобы узнать тот же текст модуля при повторном подключении + private static string GetSourceHash(string code) { - - // Convert the input string to a byte array and compute the hash. - byte[] data = md5Hash.ComputeHash(Encoding.UTF8.GetBytes(input)); - - // Create a new Stringbuilder to collect the bytes - // and create a string. - StringBuilder sBuilder = new StringBuilder(); - - // Loop through each byte of the hashed data - // and format each one as a hexadecimal string. - for (int i = 0; i < data.Length; i++) - { - sBuilder.Append(data[i].ToString("x2")); - } - - // Return the hexadecimal string. - return sBuilder.ToString(); + return Convert.ToHexString(SHA256.HashData(Encoding.UTF8.GetBytes(code))); } public void AttachByPath(ICompilerFrontend compiler, string path, string typeName, IBslProcess process) @@ -102,16 +91,12 @@ private void ThrowIfTypeExist(string typeName, SourceCode code) { if (TypeManager.IsKnownType(typeName) && _loadedModules.ContainsKey(typeName)) { - using (MD5 md5Hash = MD5.Create()) - { - string moduleCode = code.GetSourceCode(); - string hash = GetMd5Hash(md5Hash, moduleCode); - string storedHash = _fileHashes[typeName]; + string hash = GetSourceHash(code.GetSourceCode()); - StringComparer comparer = StringComparer.OrdinalIgnoreCase; - if(comparer.Compare(hash, storedHash) != 0) - throw new RuntimeException("Type «" + typeName + "» already registered"); - } + // Хеша нет у классов библиотек и у модуля, регистрацию которого откатили + StringComparer comparer = StringComparer.OrdinalIgnoreCase; + if(!_fileHashes.TryGetValue(typeName, out var storedHash) || comparer.Compare(hash, storedHash) != 0) + throw new RuntimeException("Type «" + typeName + "» already registered"); } @@ -119,38 +104,80 @@ private void ThrowIfTypeExist(string typeName, SourceCode code) private void CompileAndRegister(Type type, ICompilerFrontend compiler, string typeName, SourceCode code, IBslProcess process) { - if(_loadedModules.ContainsKey(typeName)) + string hash = GetSourceHash(code.GetSourceCode()); + + // Под блокировкой: регистрация из другого потока могла еще не закончиться или откатиться, + // а после ThrowIfTypeExist тип могли подключить с другим текстом + lock (_registrationLock) { - return; + if (IsRegisteredWithSameSource(typeName, hash)) + return; } var module = CompileModuleFromSource(compiler, code, null, process); - _loadedModules.Add(typeName, module); - using(var md5Hash = MD5.Create()) + + lock (_registrationLock) { - var hash = GetMd5Hash(md5Hash, code.GetSourceCode()); - _fileHashes.Add(typeName, hash); + // Пока модуль компилировался, этот же тип могли подключить из другого потока + if (IsRegisteredWithSameSource(typeName, hash)) + return; + + // Хеш раньше модуля: ThrowIfTypeExist читает его, когда модуль уже виден + _fileHashes[typeName] = hash; + _loadedModules[typeName] = module; + try + { + TypeManager.RegisterType(typeName, default, type); + } + catch + { + // Имя занято другим типом: без отката повторное подключение молча «удалось» бы + _loadedModules.TryRemove(typeName, out _); + _fileHashes.TryRemove(typeName, out _); + throw; + } } + } + + // true - тип уже подключен с этим же текстом, другой текст - исключение + private bool IsRegisteredWithSameSource(string typeName, string hash) + { + if (!_loadedModules.ContainsKey(typeName)) + return false; - TypeManager.RegisterType(typeName, default, type); + if (!_fileHashes.TryGetValue(typeName, out var storedHash) + || !StringComparer.OrdinalIgnoreCase.Equals(hash, storedHash)) + throw new RuntimeException("Type «" + typeName + "» already registered"); + return true; } public void RegisterTypeModule(string typeName, IExecutableModule module) { - if (_loadedModules.ContainsKey(typeName)) + lock (_registrationLock) { - var alreadyLoadedSrc = (_loadedModules[typeName]).Source.Location; - var currentSrc = module.Source.Location; + if (_loadedModules.TryGetValue(typeName, out var loadedModule)) + { + var alreadyLoadedSrc = loadedModule.Source.Location; + var currentSrc = module.Source.Location; - if(alreadyLoadedSrc != currentSrc) - throw new RuntimeException("Type «" + typeName + "» already registered"); + if(alreadyLoadedSrc != currentSrc) + throw new RuntimeException("Type «" + typeName + "» already registered"); - return; + return; + } + + _loadedModules[typeName] = module; + try + { + _engine.TypeManager.RegisterType(typeName, default, typeof(AttachedScriptsFactory)); + } + catch + { + _loadedModules.TryRemove(typeName, out _); + throw; + } } - - _loadedModules.Add(typeName, module); - _engine.TypeManager.RegisterType(typeName, default, typeof(AttachedScriptsFactory)); } private UserScriptContextInstance LoadAndCreate(ICompilerFrontend compiler, SourceCode code, diff --git a/src/ScriptEngine/Machine/DefaultTypeManager.cs b/src/ScriptEngine/Machine/DefaultTypeManager.cs index bf7b36bc3..16b6f5aa9 100644 --- a/src/ScriptEngine/Machine/DefaultTypeManager.cs +++ b/src/ScriptEngine/Machine/DefaultTypeManager.cs @@ -6,6 +6,7 @@ This Source Code Form is subject to the terms of the ----------------------------------------------------------*/ using System; +using System.Collections.Concurrent; using System.Collections.Generic; using System.Linq; using OneScript.Contexts; @@ -18,8 +19,12 @@ namespace ScriptEngine.Machine { public class DefaultTypeManager : ITypeManager { - private readonly Dictionary _knownTypesIndexes = new Dictionary(StringComparer.InvariantCultureIgnoreCase); - private readonly List _knownTypes = new List(); + // Типы регистрируются и из фоновых заданий (ПодключитьСценарий, внешние компоненты). + // Регистрация идет под блокировкой, чтение — без нее: словарь конкурентный, + // а список типов при регистрации заменяется новым массивом. + private readonly object _registrationLock = new object(); + private readonly ConcurrentDictionary _knownTypesByName = new ConcurrentDictionary(StringComparer.InvariantCultureIgnoreCase); + private volatile TypeDescriptor[] _knownTypes = Array.Empty(); private readonly TypeFactoryCache _factoryCache = new TypeFactoryCache(); private readonly ILazyTypeResolver[] _resolvers; @@ -42,9 +47,9 @@ public DefaultTypeManager(IEnumerable resolvers = null) public TypeDescriptor GetTypeByName(string name) { - if (_knownTypesIndexes.TryGetValue(name, out var index)) + if (_knownTypesByName.TryGetValue(name, out var knownType)) { - return _knownTypes[index]; + return knownType; } if (TryResolveLazily(name, out var resolvedType)) @@ -70,9 +75,8 @@ public bool TryGetType(Type frameworkType, out TypeDescriptor type) public bool TryGetType(string name, out TypeDescriptor type) { - if (_knownTypesIndexes.TryGetValue(name, out var index)) + if (_knownTypesByName.TryGetValue(name, out type)) { - type = _knownTypes[index]; return true; } @@ -87,37 +91,38 @@ public bool TryGetType(string name, out TypeDescriptor type) public TypeDescriptor RegisterType(string name, string alias, Type implementingClass) { - if (_knownTypesIndexes.ContainsKey(name)) + lock (_registrationLock) { - var td = GetTypeByName(name); - if (td.ImplementingClass != implementingClass) + if (_knownTypesByName.TryGetValue(name, out var td)) { - throw new InvalidOperationException($"Name `{name}` is already registered"); + if (td.ImplementingClass != implementingClass) + { + throw new InvalidOperationException($"Name `{name}` is already registered"); + } + + return td; } - return td; - } - else - { var typeDesc = new TypeDescriptor(implementingClass, name, alias); RegisterTypeInternal(typeDesc); return typeDesc; } - } public void RegisterType(TypeDescriptor typeDescriptor) { - if (_knownTypesIndexes.TryGetValue(typeDescriptor.Name, out var index)) + lock (_registrationLock) { - var knownType = _knownTypes[index]; - if (knownType != typeDescriptor) - throw new InvalidOperationException($"Type {typeDescriptor} already registered"); - - return; + if (_knownTypesByName.TryGetValue(typeDescriptor.Name, out var knownType)) + { + if (knownType != typeDescriptor) + throw new InvalidOperationException($"Type {typeDescriptor} already registered"); + + return; + } + + RegisterTypeInternal(typeDescriptor); } - - RegisterTypeInternal(typeDescriptor); } public ITypeFactory GetFactoryFor(TypeDescriptor type) @@ -127,12 +132,16 @@ public ITypeFactory GetFactoryFor(TypeDescriptor type) private void RegisterTypeInternal(TypeDescriptor td) { - var nextListId = _knownTypes.Count; - _knownTypesIndexes.Add(td.Name, nextListId); + // Сначала список: тип, найденный по имени, уже есть и в нем + var knownTypes = _knownTypes; + var newKnownTypes = new TypeDescriptor[knownTypes.Length + 1]; + knownTypes.CopyTo(newKnownTypes, 0); + newKnownTypes[knownTypes.Length] = td; + _knownTypes = newKnownTypes; + + _knownTypesByName[td.Name] = td; if (!string.IsNullOrWhiteSpace(td.Alias) && td.Alias != td.Name) - _knownTypesIndexes[td.Alias] = nextListId; - - _knownTypes.Add(td); + _knownTypesByName[td.Alias] = td; } private bool TryResolveLazily(string name, out TypeDescriptor type) diff --git a/src/Tests/OneScript.Core.Tests/TestTypes_Registration.cs b/src/Tests/OneScript.Core.Tests/TestTypes_Registration.cs index 65861a9c9..205c2a92c 100644 --- a/src/Tests/OneScript.Core.Tests/TestTypes_Registration.cs +++ b/src/Tests/OneScript.Core.Tests/TestTypes_Registration.cs @@ -11,9 +11,11 @@ This Source Code Form is subject to the terms of the using FluentAssertions; using Moq; using OneScript.Contexts; +using OneScript.Exceptions; using OneScript.StandardLibrary.Collections; using OneScript.Types; using ScriptEngine; +using ScriptEngine.Hosting; using ScriptEngine.Machine; using ScriptEngine.Machine.Contexts; using ScriptEngine.Types; @@ -98,6 +100,39 @@ public void TypeEqualityForNull() Assert.True(type1 == type2); // operator== Assert.False(type1 != type2); // operator != } - + + [Fact] + public void AttachingScriptUnderBuiltInTypeNameFailsEveryTime() + { + var engine = DefaultEngineBuilder.Create() + .SetDefaultOptions() + .Build(); + engine.Initialize(); + + for (var i = 0; i < 2; i++) + { + Action attach = () => engine.AttachedScriptsFactory.AttachFromString( + engine.GetCompilerService(), "Перем А;", "Строка", engine.NewProcess()); + attach.Should().Throw(); + } + } + + [Fact] + public void AttachingScriptUnderLibraryClassNameIsRejected() + { + var engine = DefaultEngineBuilder.Create() + .SetDefaultOptions() + .Build(); + engine.Initialize(); + var libraryClass = engine.AttachedScriptsFactory.CompileModuleFromSource( + engine.GetCompilerService(), engine.Loader.FromString("Перем А;"), null, engine.NewProcess()); + engine.AttachedScriptsFactory.RegisterTypeModule("КлассБиблиотеки", libraryClass); + + Action attach = () => engine.AttachedScriptsFactory.AttachFromString( + engine.GetCompilerService(), "Перем Б;", "КлассБиблиотеки", engine.NewProcess()); + + attach.Should().Throw().WithMessage("*already registered*"); + } + } } \ No newline at end of file diff --git a/src/Tests/OneScript.Core.Tests/TypeRegistrationThreadSafetyTests.cs b/src/Tests/OneScript.Core.Tests/TypeRegistrationThreadSafetyTests.cs new file mode 100644 index 000000000..fa2794b74 --- /dev/null +++ b/src/Tests/OneScript.Core.Tests/TypeRegistrationThreadSafetyTests.cs @@ -0,0 +1,137 @@ +/*---------------------------------------------------------- +This Source Code Form is subject to the terms of the +Mozilla Public License, v.2.0. If a copy of the MPL +was not distributed with this file, You can obtain one +at http://mozilla.org/MPL/2.0/. +----------------------------------------------------------*/ + +using System; +using System.Collections.Concurrent; +using System.Linq; +using System.Threading; +using FluentAssertions; +using OneScript.Types; +using OneScript.Values; +using ScriptEngine.Hosting; +using ScriptEngine.Machine; +using ScriptEngine.Machine.Contexts; +using Xunit; + +namespace OneScript.Core.Tests +{ + public class TypeRegistrationThreadSafetyTests + { + [Fact] + public void TypesRegisteredFromManyThreadsAreFoundByName() + { + const int threadsCount = 16; + const int typesPerThread = 500; + var typeManager = new DefaultTypeManager(); + var builtInCount = typeManager.RegisteredTypes().Count; + + RunInParallel(threadsCount, thread => + { + for (var i = 0; i < typesPerThread; i++) + { + var registered = typeManager.RegisterType($"Тип{thread}_{i}", $"Type{thread}_{i}", typeof(BslValue)); + + typeManager.GetTypeByName($"Тип{thread}_{i}").Should().BeSameAs(registered); + typeManager.GetTypeByName($"Type{thread}_{i}").Should().BeSameAs(registered); + // Другие потоки в это время регистрируют свои типы + typeManager.RegisteredTypes().Should().Contain(registered); + } + }); + + typeManager.RegisteredTypes().Should().HaveCount(builtInCount + threadsCount * typesPerThread); + for (var thread = 0; thread < threadsCount; thread++) + { + for (var i = 0; i < typesPerThread; i++) + { + typeManager.GetTypeByName($"Type{thread}_{i}").Name.Should().Be($"Тип{thread}_{i}"); + } + } + } + + [Fact] + public void SameTypeRegisteredConcurrentlyGetsOneDescriptor() + { + const int threadsCount = 8; + var typeManager = new DefaultTypeManager(); + + for (var round = 0; round < 200; round++) + { + var name = $"Общий{round}"; + var results = new ConcurrentBag(); + + RunInParallel(threadsCount, _ => results.Add(typeManager.RegisterType(name, default, typeof(BslValue)))); + + results.Distinct().Should().ContainSingle(name); + typeManager.RegisteredTypes().Count(x => x.Name == name).Should().Be(1, name); + } + } + + [Fact] + public void ScriptsAttachedFromManyThreadsAreRegisteredOnce() + { + const int threadsCount = 16; + const int classesPerThread = 30; + var engine = DefaultEngineBuilder.Create() + .SetDefaultOptions() + .Build(); + engine.Initialize(); + + RunInParallel(threadsCount, thread => + { + var process = engine.NewProcess(); + for (var i = 0; i < classesPerThread; i++) + { + // Общий класс подключают все потоки, пока каждый подключает и свои классы + engine.AttachedScriptsFactory.AttachFromString(engine.GetCompilerService(), + "Перем Общая;", "ОбщийКласс", process); + engine.AttachedScriptsFactory.AttachFromString(engine.GetCompilerService(), + $"Перем Номер{i};", $"Класс{thread}_{i}", process); + } + }); + + engine.TypeManager.RegisteredTypes().Count(x => x.Name == "ОбщийКласс").Should().Be(1); + for (var thread = 0; thread < threadsCount; thread++) + { + for (var i = 0; i < classesPerThread; i++) + { + engine.TypeManager.GetTypeByName($"Класс{thread}_{i}").ImplementingClass + .Should().Be(typeof(AttachedScriptsFactory)); + } + } + } + + private static void RunInParallel(int threadsCount, Action action) + { + var exceptions = new ConcurrentQueue(); + using var barrier = new Barrier(threadsCount); + var threads = Enumerable.Range(0, threadsCount).Select(number => new Thread(() => + { + try + { + barrier.SignalAndWait(); + action(number); + } + catch (Exception e) + { + exceptions.Enqueue(e); + } + }) { IsBackground = true }).ToArray(); + + foreach (var thread in threads) + { + thread.Start(); + } + foreach (var thread in threads) + { + // Испорченный словарь может зациклить поток + thread.Join(TimeSpan.FromSeconds(60)).Should().BeTrue("поток должен завершиться"); + } + + exceptions.Should().BeEmpty(); + } + } +}