diff --git a/.github/import_generation.txt b/.github/import_generation.txt index e373ee695f6..82cced27d7b 100644 --- a/.github/import_generation.txt +++ b/.github/import_generation.txt @@ -1 +1 @@ -50 +51 diff --git a/.github/last_commit.txt b/.github/last_commit.txt index d081fe0e422..575ceed5215 100644 --- a/.github/last_commit.txt +++ b/.github/last_commit.txt @@ -1 +1 @@ -4b757a14e7269921b351fec8dff5af2e088a89ae +c7fb4a28d8e03aaa9e42859b6f7c199c6e7ef203 diff --git a/CHANGELOG.md b/CHANGELOG.md index 1d75b8b78c0..d49e63c66e4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,19 @@ +## v3.24.0 + +* Added `NValueHelpers::Embedding` to create FloatVector `Bytes` query parameters from numeric ranges. + +* Topic and PersQueue writers now respect the driver's outbound gRPC message size limit when batching writes. + +* Added a draft UDF client (`client/draft/ydb_udf.h`) with manifest-based uploads, separate module type/code kind, per-platform compile state and optional timestamps, and incremental `UploadModuleFromFile` on a dedicated I/O executor. Upload futures include the final gRPC status. + +* Added an optional S3 object key prefix to TTL eviction settings for column tables. + +* Added `TTopicClient::ResetOffset` / `TResetOffsetSettings` to rewind a consumer's committed offsets on all topic partitions. + +* Added OIDC/OAuth authentication via `NOidc::CreateOidcProviderFactory`, supporting static access tokens, Client Credentials Grant, and Device Authorization Grant, with token refresh and interfaces for token caching and interactive sign-in. + +* Added `Float16` and `BFloat16` vector index types. + ## v3.23.0 * Added optional `TRetryOperationSettings::StopToken` for cooperative cancellation between retry attempts. diff --git a/examples/auth/oidc/main.cpp b/examples/auth/oidc/main.cpp new file mode 100644 index 00000000000..8dda22499ea --- /dev/null +++ b/examples/auth/oidc/main.cpp @@ -0,0 +1,93 @@ +#include +#include + +#include + +#include + +class TConsoleAcceptor final: public NYdb::NOidc::IAuthAcceptor { +public: + void Accept(const NYdb::NOidc::TDeviceAuthInfo& info) override; + +private: + static void Print(const std::string& value); +}; + +class TMemoryTokenCacher final: public NYdb::NOidc::ITokenCacher { +public: + std::optional Read() const override; + void Write(const NYdb::NOidc::TTokenCache& tokens) override; + +private: + mutable TMutex Mutex; + std::optional Tokens; +}; + +void TConsoleAcceptor::Accept(const NYdb::NOidc::TDeviceAuthInfo& info) { + static TMutex outputMutex; + with_lock (outputMutex) { + std::cerr << "Open "; + Print(info.VerificationUrl); + std::cerr << " and enter code "; + Print(info.UserCode); + std::cerr << std::endl; + if (info.VerificationUrlComplete.has_value()) { + std::cerr << "Or open "; + Print(*info.VerificationUrlComplete); + std::cerr << std::endl; + } + } +} + +void TConsoleAcceptor::Print(const std::string& value) { + for (const unsigned char ch : value) { + std::cerr << (ch >= 0x20 && ch < 0x7f ? static_cast(ch) : '?'); + } +} + +std::optional TMemoryTokenCacher::Read() const { + with_lock (Mutex) { + return Tokens; + } +} + +void TMemoryTokenCacher::Write(const NYdb::NOidc::TTokenCache& tokens) { + with_lock (Mutex) { + Tokens = tokens; + } +} + +int main(int argc, char** argv) { + if (argc != 5) { + std::cerr << "Usage: oidc " << std::endl; + return 1; + } + try { + NYdb::NOidc::TOidcConfig oidcConfig{ + .Issuer = argv[3], + .FlowConfig = NYdb::NOidc::TDeviceOidcConfig{ + .ClientId = argv[4], + .Scopes = {"openid", "user-context"}, + }, + }; + oidcConfig + .Cacher(std::make_shared()) + .Acceptor(std::make_shared()); + + auto config = NYdb::TDriverConfig(argv[1]) + .SetDatabase(argv[2]) + .SetDiscoveryMode(NYdb::EDiscoveryMode::Async) + .SetCredentialsProviderFactory(NYdb::NOidc::CreateOidcProviderFactory(oidcConfig)); + NYdb::TDriver driver(config); + NYdb::NQuery::TQueryClient client(driver); + auto result = client.ExecuteQuery("SELECT 1", NYdb::NQuery::TTxControl::NoTx()).ExtractValueSync(); + if (!result.IsSuccess()) { + std::cerr << result.GetIssues().ToString() << std::endl; + return 1; + } + std::cout << "Query succeeded" << std::endl; + } catch (const std::exception& error) { + std::cerr << error.what() << std::endl; + return 1; + } +} diff --git a/examples/vector_index/README.md b/examples/vector_index/README.md index 244d1fcf525..e568cabeb80 100644 --- a/examples/vector_index/README.md +++ b/examples/vector_index/README.md @@ -42,7 +42,7 @@ It uses scalar quantization to speedup ANN search: #### Create Flat Bit Index -It creates index table, with two columns: `primary_key` and `embedding`, `embedding` will be trasformed from `` `` +It creates index table, with two columns: `primary_key` and `embedding`, `embedding` will be transformed from `
` `` ``` ./vector_index --endpoint= --database= --command=RecreateIndex --table=
--index_type=flat --index_quantizer=bit --primary_key= --embedding= --distance=CosineDistance --rows=490000 --top_k=0 --data="" --target="" diff --git a/include/ydb-cpp-sdk/client/query/client.h b/include/ydb-cpp-sdk/client/query/client.h index e9bc56bdb61..9d743f3aa4c 100644 --- a/include/ydb-cpp-sdk/client/query/client.h +++ b/include/ydb-cpp-sdk/client/query/client.h @@ -31,7 +31,9 @@ namespace NYdb::inline V3 { namespace NYdb::inline V3::NQuery { +//! Request settings for obtaining a query session. struct TCreateSessionSettings : public TSimpleRequestSettings { + //! Constructs settings with a five-second client timeout. TCreateSessionSettings() { ClientTimeout(TDuration::Seconds(5)); } @@ -40,37 +42,39 @@ struct TCreateSessionSettings : public TSimpleRequestSettings; using TRetryOperationSettings = NYdb::NRetry::TRetryOperationSettings; +//! Settings for the query session pool. struct TSessionPoolSettings { using TSelf = TSessionPoolSettings; - // Max number of sessions client can get from session pool + //! Sets the maximum number of sessions that may be acquired from the pool; defaults to 50. FLUENT_SETTING_DEFAULT(uint32_t, MaxActiveSessions, 50); - // Max time session to be in idle state before closing + //! Sets how long an idle session may remain in the pool before closing; defaults to one minute. FLUENT_SETTING_DEFAULT(TDuration, CloseIdleThreshold, TDuration::Minutes(1)); - // Min number of session in session pool. - // Sessions will not be closed by CloseIdleThreshold if the number of sessions less then this limit. + //! Sets the minimum pool size protected from idle eviction; defaults to 10 sessions. FLUENT_SETTING_DEFAULT(uint32_t, MinPoolSize, 10); - // Create session in the background even after client timeout. - // This is useful for applications with short session timeouts. + //! Keeps creating a session in the background after the caller's timeout; disabled by default. FLUENT_SETTING_DEFAULT(bool, UseDeferredSessionCreation, false); }; +//! Settings shared by all operations performed through TQueryClient. struct TClientSettings : public TCommonClientSettingsBase { using TSessionPoolSettings = TSessionPoolSettings; using TSelf = TClientSettings; + //! Configures the query session pool. FLUENT_SETTING(TSessionPoolSettings, SessionPoolSettings); - // Optional pool name surfaced through the OTel tag - // ydb.query.session.pool.name. When empty the default - // "@" is used. + //! Sets the value of the ydb.query.session.pool.name OpenTelemetry tag. + //! An empty value uses "@". FLUENT_SETTING(std::string, PoolName); + //! Sets the default retry policy for retry-capable query client operations. FLUENT_SETTING_DEFAULT(TRetryOperationSettings, RetrySettings, TRetryOperationSettings()); }; +//! Client for executing queries and scripts through YDB Query Service. class TQueryClient { friend class TSession; template @@ -90,60 +94,89 @@ class TQueryClient { using TAsyncCreateSessionResult = TAsyncCreateSessionResult; public: + //! Constructs a query client that uses the supplied driver and settings. TQueryClient(const TDriver& driver, const TClientSettings& settings = TClientSettings()); + //! Executes a query without parameters or an explicit session and buffers all result sets. + //! Multi-step interactive transactions must be executed through TSession. TAsyncExecuteQueryResult ExecuteQuery(const std::string& query, const TTxControl& txControl, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Executes a parameterized query without an explicit session and buffers all result sets. + //! Multi-step interactive transactions must be executed through TSession. TAsyncExecuteQueryResult ExecuteQuery(const std::string& query, const TTxControl& txControl, const TParams& params, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Starts streaming a query without parameters or an explicit session. + //! Multi-step interactive transactions must be executed through TSession. TAsyncExecuteQueryIterator StreamExecuteQuery(const std::string& query, const TTxControl& txControl, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Starts streaming a parameterized query without an explicit session. + //! Multi-step interactive transactions must be executed through TSession. TAsyncExecuteQueryIterator StreamExecuteQuery(const std::string& query, const TTxControl& txControl, const TParams& params, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Runs an asynchronous result-returning callback with a pooled session and retries retryable failures. + //! The callback may be invoked more than once and must honor the idempotency configured in settings. TAsyncExecuteQueryResult RetryQuery(TQueryResultFunc&& queryFunc, TRetryOperationSettings settings = TRetryOperationSettings()); + //! Runs an asynchronous status-returning callback with a pooled session and retries retryable failures. + //! The callback may be invoked more than once and must honor the idempotency configured in settings. TAsyncStatus RetryQuery(TQueryFunc&& queryFunc, TRetryOperationSettings settings = TRetryOperationSettings()); + //! Runs an asynchronous callback without a session and retries retryable failures. + //! The callback may be invoked more than once and must honor the idempotency configured in settings. TAsyncStatus RetryQuery(TQueryWithoutSessionFunc&& queryFunc, TRetryOperationSettings settings = TRetryOperationSettings()); + //! Runs a synchronous callback with a pooled session and retries retryable failures. + //! The callback may be invoked more than once and must honor the idempotency configured in settings. TStatus RetryQuerySync(const TQuerySyncFunc& queryFunc, TRetryOperationSettings settings = TRetryOperationSettings()); + //! Runs a synchronous callback without a session and retries retryable failures. + //! The callback may be invoked more than once and must honor the idempotency configured in settings. TStatus RetryQuerySync(const TQueryWithoutSessionSyncFunc& queryFunc, TRetryOperationSettings settings = TRetryOperationSettings()); + //! Executes a query with retries for up to timeout. isIdempotent specifies whether the operation is + //! idempotent and can therefore be retried safely. TAsyncExecuteQueryResult RetryQuery(const std::string& query, const TTxControl& txControl, - TDuration timeout, bool isIndempotent); + TDuration timeout, bool isIdempotent); + //! Starts asynchronous execution of a script without parameters. NThreading::TFuture ExecuteScript(const std::string& script, const TExecuteScriptSettings& settings = TExecuteScriptSettings(), const std::optional& retrySettings = std::nullopt); + //! Starts asynchronous execution of a parameterized script. NThreading::TFuture ExecuteScript(const std::string& script, const TParams& params, const TExecuteScriptSettings& settings = TExecuteScriptSettings(), const std::optional& retrySettings = std::nullopt); + //! Fetches one page of a completed script result set. TAsyncFetchScriptResultsResult FetchScriptResults(const NKikimr::NOperationId::TOperationId& operationId, int64_t resultSetIndex, const TFetchScriptResultsSettings& settings = TFetchScriptResultsSettings(), const std::optional& retrySettings = std::nullopt); + //! Acquires a session from the internal session pool or creates a new one. TAsyncCreateSessionResult GetSession(const TCreateSessionSettings& settings = TCreateSessionSettings()); + //! Explicitly deletes the server-side session identified by sessionId. TAsyncStatus DeleteSession(const std::string& sessionId, const TDeleteSessionSettings& settings = TDeleteSessionSettings()); - //! Returns number of active sessions given via session pool + //! Returns the number of sessions currently acquired from the session pool. int64_t GetActiveSessionCount() const; - //! Returns the maximum number of sessions in session pool + //! Returns the maximum number of sessions that may be acquired from the session pool. int64_t GetActiveSessionsLimit() const; - //! Returns the size of session pool + //! Returns the number of idle sessions currently available in the session pool. int64_t GetCurrentPoolSize() const; - // Internal: used by retry wrappers to suppress nested retries. + //! Returns whether the current thread is already inside a query retry operation. + //! This method is intended for SDK retry wrappers. bool GetInRetryOperationContext() const; + //! Marks whether the current thread is inside a query retry operation. + //! This method is intended for SDK retry wrappers. void SetInRetryOperationContext(bool value); private: @@ -151,28 +184,36 @@ class TQueryClient { std::shared_ptr Impl_; }; +//! A Query Service session used for interactive transactions and session-bound queries. class TSession { friend class TQueryClient; friend class TTransaction; friend class TExecuteQueryIterator; friend class NRetry::TRetryDeadlineHelper; public: + //! Returns the server-side session identifier. const std::string& GetId() const; + //! Returns the deadline propagated to this session by a retry operation, when present. const std::optional& GetPropagatedDeadline() const; + //! Executes a query without parameters in this session and buffers all result sets. TAsyncExecuteQueryResult ExecuteQuery(const std::string& query, const TTxControl& txControl, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Executes a parameterized query in this session and buffers all result sets. TAsyncExecuteQueryResult ExecuteQuery(const std::string& query, const TTxControl& txControl, const TParams& params, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Starts streaming a query without parameters in this session. TAsyncExecuteQueryIterator StreamExecuteQuery(const std::string& query, const TTxControl& txControl, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Starts streaming a parameterized query in this session. TAsyncExecuteQueryIterator StreamExecuteQuery(const std::string& query, const TTxControl& txControl, const TParams& params, const TExecuteQuerySettings& settings = TExecuteQuerySettings()); + //! Begins an interactive transaction in this session. TAsyncBeginTransactionResult BeginTransaction(const TTxSettings& txSettings, const TBeginTxSettings& settings = TBeginTxSettings()); @@ -188,30 +229,40 @@ class TSession { std::shared_ptr SessionImpl_; }; +//! Result of acquiring or creating a query session. class TCreateSessionResult: public TStatus { friend class TSession::TImpl; public: + //! Constructs a session result from a status and session. TCreateSessionResult(TStatus&& status, TSession&& session); + //! Returns the session, throwing if the operation status is not successful. TSession GetSession() const; private: TSession Session_; }; +//! An interactive Query Service transaction bound to a TSession. class TTransaction : public TTransactionBase { friend class TQueryClient; friend class TExecuteQueryIterator::TReaderImpl; friend class TExecQueryImpl; public: + //! Returns whether this transaction has a non-empty server-side identifier. bool IsActive() const; + //! Runs precommit callbacks and asynchronously commits the transaction. TAsyncCommitTransactionResult Commit(const TCommitTxSettings& settings = TCommitTxSettings()); + //! Asynchronously rolls back the transaction and runs failure callbacks. TAsyncStatus Rollback(const TRollbackTxSettings& settings = TRollbackTxSettings()); + //! Returns the session to which this transaction is bound. TSession GetSession() const; + //! Registers a callback that is run before the transaction is committed. void AddPrecommitCallback(TPrecommitTransactionCallback cb) override; + //! Registers a callback that is run after commit failure or rollback. void AddOnFailureCallback(TOnFailureTransactionCallback cb) override; private: @@ -225,6 +276,7 @@ class TTransaction : public TTransactionBase { std::shared_ptr TransactionImpl_; }; +//! Describes how a query participates in a transaction. class TTxControl { friend class TExecQueryImpl; friend class TExecQueryInternal; @@ -232,27 +284,33 @@ class TTxControl { public: using TSelf = TTxControl; + //! Continues an existing interactive transaction. static TTxControl Tx(const TTransaction& tx) { return TTxControl(tx); } + //! Continues a transaction by its identifier. + //! Prefer Tx(const TTransaction&) so that the owning session is retained. [[deprecated("This is bug-provoking API. Use TTxControl::Tx(TTransaction) instead. " "This constructor will be removed in upcomming release")]] static TTxControl Tx(const std::string& txId) { return TTxControl(txId); } + //! Begins a transaction with the supplied transaction settings. static TTxControl BeginTx(const TTxSettings& settings = TTxSettings()) { return TTxControl(settings); } - // Do not explicitly set the transaction mode. YDB determines the behavior automatically + //! Leaves transaction handling to YDB's implicit transaction rules. static TTxControl NoTx() { return TTxControl(); } + //! Requests that the selected or newly started transaction be committed after the query. FLUENT_SETTING_FLAG(CommitTx); + //! Returns whether this control selects or starts an explicit transaction. bool HasTx() const { return !std::holds_alternative(Tx_); } private: @@ -270,31 +328,45 @@ class TTxControl { const std::variant Tx_; }; +//! Result of beginning an interactive transaction. class TBeginTransactionResult : public TStatus { public: + //! Constructs a begin-transaction result from a status and transaction. TBeginTransactionResult(TStatus&& status, TTransaction transaction); + //! Returns the transaction, throwing if the operation status is not successful. const TTransaction& GetTransaction() const; private: TTransaction Transaction_; }; +//! One part of a streaming query response. class TExecuteQueryPart : public TStreamPartStatus { public: + //! Returns whether this response part contains a result set fragment. bool HasResultSet() const { return ResultSet_.has_value(); } + //! Returns the zero-based index of the result set. HasResultSet() must be true. uint64_t GetResultSetIndex() const { return ResultSetIndex_; } + //! Returns the result set fragment. HasResultSet() must be true. const TResultSet& GetResultSet() const { return *ResultSet_; } + //! Moves the result set fragment out of this response part. HasResultSet() must be true. TResultSet ExtractResultSet() { return std::move(*ResultSet_); } + //! Returns whether this response part contains execution statistics. bool HasStats() const { return Stats_.has_value(); } + //! Returns the execution statistics carried by this response part, when present. const std::optional& GetStats() const { return Stats_; } + //! Returns the execution statistics. HasStats() must be true. TExecStats ExtractStats() const { return std::move(*Stats_); } + //! Returns a transaction started by the query, when it was not committed in the same request. const std::optional& GetTransaction() const { return Transaction_; } + //! Returns the commit timestamp when the query committed a transaction with write effects. const std::optional& GetCommitTimestamp() const { return CommitTimestamp_; } + //! Constructs a response part without a result set fragment. TExecuteQueryPart(TStatus&& status, std::optional&& queryStats, std::optional&& tx, std::optional&& commitTimestamp = {}) : TStreamPartStatus(std::move(status)) @@ -303,6 +375,7 @@ class TExecuteQueryPart : public TStreamPartStatus { , CommitTimestamp_(std::move(commitTimestamp)) {} + //! Constructs a response part containing a result set fragment. TExecuteQueryPart(TStatus&& status, TResultSet&& resultSet, int64_t resultSetIndex, std::optional&& queryStats, std::optional&& tx, std::optional&& commitTimestamp = {}) @@ -322,22 +395,31 @@ class TExecuteQueryPart : public TStreamPartStatus { std::optional CommitTimestamp_; }; +//! Buffered result of a query execution. class TExecuteQueryResult : public TStatus { public: + //! Returns all result sets in statement order. const std::vector& GetResultSets() const; + //! Returns a copy of the result set at resultIndex, throwing when the index is out of range. TResultSet GetResultSet(size_t resultIndex) const; + //! Returns a parser for the result set at resultIndex, throwing when the index is out of range. TResultSetParser GetResultSetParser(size_t resultIndex) const; + //! Returns execution statistics when statistics collection was enabled. const std::optional& GetStats() const { return Stats_; } + //! Returns a transaction started by the query, when it was not committed in the same request. std::optional GetTransaction() const {return Transaction_; } + //! Returns the commit timestamp when the query committed a transaction with write effects. const std::optional& GetCommitTimestamp() const { return CommitTimestamp_; } + //! Constructs a query result that contains only an operation status. TExecuteQueryResult(TStatus&& status) : TStatus(std::move(status)) {} + //! Constructs a query result from its status, result sets, statistics, and transaction metadata. TExecuteQueryResult(TStatus&& status, std::vector&& resultSets, std::optional&& stats, std::optional&& tx, std::optional&& commitTimestamp = {}) diff --git a/include/ydb-cpp-sdk/client/query/query.h b/include/ydb-cpp-sdk/client/query/query.h index c4acd8c1f72..65e714739b7 100644 --- a/include/ydb-cpp-sdk/client/query/query.h +++ b/include/ydb-cpp-sdk/client/query/query.h @@ -18,74 +18,116 @@ namespace NYdb::inline V3::NQuery { using TRetryOperationSettings = NYdb::NRetry::TRetryOperationSettings; +//! Query text syntax accepted by Query Service. enum class ESyntax { + //! Let the server choose the query syntax. Unspecified = 0, - YqlV1 = 1, // YQL - Pg = 2, // PostgresQL + //! YQL version 1 syntax. + YqlV1 = 1, + //! PostgreSQL-compatible syntax. + Pg = 2, }; +//! Controls how far the server processes a query. enum class EExecMode { + //! Let the server choose the execution mode. Unspecified = 0, + //! Parse the query without validating or executing it. Parse = 10, + //! Parse and validate the query without executing it. Validate = 20, + //! Build an execution plan without executing the query. Explain = 30, + //! Execute the query. Execute = 50, }; +//! Controls the amount of execution statistics returned by the server. enum class EStatsMode { + //! Let the server choose the statistics mode. Unspecified = 0, + //! Do not collect execution statistics. None = 10, + //! Collect aggregated table access statistics. Basic = 20, + //! Add execution statistics and the query plan to basic statistics. Full = 30, + //! Collect detailed task and channel statistics. Profile = 40, }; +//! Controls how often a result set schema is included in a response stream. enum class ESchemaInclusionMode { + //! Use the server default, which is equivalent to Always. Unspecified = 0, + //! Include the schema in every result set part. Always = 1, + //! Include the schema only in the first part of each result set. FirstOnly = 2, }; +//! Parses a lowercase statistics mode name, or returns std::nullopt for an unknown name. std::optional ParseStatsMode(std::string_view statsMode); +//! Returns the lowercase name of a statistics mode. std::string_view StatsModeToString(const EStatsMode statsMode); +//! Current state of an asynchronously executed script. enum class EExecStatus { + //! The execution state is not specified. Unspecified = 0, + //! The script is being prepared for execution. Starting = 10, + //! The script is running. Running = 15, + //! The script was aborted by the server. Aborted = 20, + //! The script was canceled. Canceled = 30, + //! The script completed successfully. Completed = 40, + //! The script execution failed. Failed = 50, }; +//! Apache Arrow result format settings. struct TArrowFormatSettings { using TSelf = TArrowFormatSettings; + //! Compression settings for Arrow record batches. struct TCompressionCodec { using TSelf = TCompressionCodec; + //! Supported Arrow record batch compression codecs. enum class EType { + //! Use the server default, which is equivalent to None. Unspecified = 0, + //! Do not compress record batches. None = 1, + //! Use Zstandard compression. Zstd = 2, + //! Use LZ4 frame compression. Lz4Frame = 3, }; + //! Selects the compression codec; defaults to the server choice. FLUENT_SETTING_DEFAULT(EType, Type, EType::Unspecified); + //! Sets the codec-specific compression level; the codec default is used when unset. FLUENT_SETTING_OPTIONAL(int32_t, Level); }; + //! Sets compression options for Arrow record batches. FLUENT_SETTING_OPTIONAL(TCompressionCodec, CompressionCodec); }; using TAsyncExecuteQueryPart = NThreading::TFuture; +//! Asynchronous reader for a streaming query response. class TExecuteQueryIterator : public TStatus { friend class TExecQueryImpl; public: class TReaderImpl; + //! Asynchronously reads the next response part. Read until the returned part reports EOS(). TAsyncExecuteQueryPart ReadNext(); private: @@ -106,35 +148,57 @@ class TExecuteQueryIterator : public TStatus { using TAsyncExecuteQueryIterator = NThreading::TFuture; +//! Settings for executing or streaming a query. struct TExecuteQuerySettings : public TRequestSettings { + //! Limits one streamed result part to the specified number of bytes. FLUENT_SETTING_OPTIONAL(uint32_t, OutputChunkMaxSize); + //! Selects the query syntax; defaults to YQL version 1. FLUENT_SETTING_DEFAULT(ESyntax, Syntax, ESyntax::YqlV1); + //! Selects the execution mode; defaults to executing the query. FLUENT_SETTING_DEFAULT(EExecMode, ExecMode, EExecMode::Execute); + //! Selects the statistics detail level; statistics are disabled by default. FLUENT_SETTING_DEFAULT(EStatsMode, StatsMode, EStatsMode::None); + //! Enables collection of affected row statistics; disabled by default and may add overhead. FLUENT_SETTING_DEFAULT(bool, CollectAffectedRows, false); + //! Allows parts of different result sets to be interleaved in the response stream. FLUENT_SETTING_OPTIONAL(bool, ConcurrentResultSets); + //! Selects the workload manager resource pool used to execute the query. FLUENT_SETTING(std::string, ResourcePool); + //! Requests periodic statistics at this interval while statistics collection is enabled. FLUENT_SETTING_OPTIONAL(std::chrono::milliseconds, StatsCollectPeriod); + //! Selects how often result set schemas are included; defaults to the server behavior. FLUENT_SETTING_DEFAULT(ESchemaInclusionMode, SchemaInclusionMode, ESchemaInclusionMode::Unspecified); + //! Selects the result set representation; defaults to the server's value format. FLUENT_SETTING_DEFAULT(TResultSet::EFormat, Format, TResultSet::EFormat::Unspecified); + //! Sets Arrow-specific options used when Format is the Arrow representation. FLUENT_SETTING_OPTIONAL(TArrowFormatSettings, ArrowFormatSettings); + //! Sets a per-request retry policy, overriding the client retry settings when present. FLUENT_SETTING_OPTIONAL(TRetryOperationSettings, RetrySettings); }; +//! Request settings for beginning a transaction. struct TBeginTxSettings : public TRequestSettings {}; +//! Request settings for committing a transaction. struct TCommitTxSettings : public TRequestSettings {}; +//! Request settings for rolling back a transaction. struct TRollbackTxSettings : public TRequestSettings {}; +//! Request settings for deleting a session. struct TDeleteSessionSettings : public TRequestSettings { + //! Sets a per-request retry policy, overriding the client retry settings when present. FLUENT_SETTING_OPTIONAL(TRetryOperationSettings, RetrySettings); }; +//! Result of committing an interactive transaction. class TCommitTransactionResult : public TStatus { public: + //! Constructs a commit result without a commit timestamp. TCommitTransactionResult(TStatus&& status); + //! Constructs a commit result with an optional commit timestamp. TCommitTransactionResult(TStatus&& status, std::optional&& commitTimestamp); + //! Returns the commit timestamp when the successful transaction produced write effects. const std::optional& GetCommitTimestamp() const { return CommitTimestamp_; } private: @@ -144,64 +208,94 @@ class TCommitTransactionResult : public TStatus { using TAsyncBeginTransactionResult = NThreading::TFuture; using TAsyncCommitTransactionResult = NThreading::TFuture; +//! Settings for starting asynchronous script execution. struct TExecuteScriptSettings : public TOperationRequestSettings { + //! Selects the script syntax; defaults to YQL version 1. FLUENT_SETTING_DEFAULT(ESyntax, Syntax, ESyntax::YqlV1); + //! Selects the execution mode; defaults to executing the script. FLUENT_SETTING_DEFAULT(EExecMode, ExecMode, EExecMode::Execute); + //! Selects the statistics detail level; statistics are disabled by default. FLUENT_SETTING_DEFAULT(EStatsMode, StatsMode, EStatsMode::None); + //! Sets how long completed script results remain available for fetching. FLUENT_SETTING(TDuration, ResultsTtl); + //! Selects the workload manager resource pool used to execute the script. FLUENT_SETTING(std::string, ResourcePool); + //! Sets a per-request retry policy, overriding the client retry settings when present. FLUENT_SETTING_OPTIONAL(TRetryOperationSettings, RetrySettings); }; +//! Query text together with its syntax. class TQueryContent { public: + //! Constructs empty query content with unspecified syntax. TQueryContent() = default; + //! Constructs query content from text and syntax. TQueryContent(const std::string& text, ESyntax syntax) : Text(text) , Syntax(syntax) {} + //! Query text. std::string Text; + //! Syntax used by the query text. ESyntax Syntax = ESyntax::Unspecified; }; +//! Metadata describing one script result set. class TResultSetMeta { public: + //! Constructs empty result set metadata. TResultSetMeta() = default; + //! Constructs metadata by copying the result set columns. explicit TResultSetMeta(const std::vector& columns, uint64_t rowsCount = 0, bool finished = false) : Columns(columns) , RowsCount(rowsCount) , Finished(finished) {} + //! Constructs metadata by taking ownership of the result set columns. explicit TResultSetMeta(std::vector&& columns, uint64_t rowsCount = 0, bool finished = false) : Columns(std::move(columns)) , RowsCount(rowsCount) , Finished(finished) {} + //! Result set column descriptions. std::vector Columns; + //! Number of rows currently available in the result set. uint64_t RowsCount = 0; + //! Whether the result set is complete. bool Finished = false; }; +//! Long-running operation returned by ExecuteScript(). class TScriptExecutionOperation : public TOperation { public: + //! Script execution metadata reported by the server. struct TMetadata { + //! Server-side script execution identifier. std::string ExecutionId; + //! Current script execution state. EExecStatus ExecStatus = EExecStatus::Unspecified; + //! Execution mode used for the script. EExecMode ExecMode = EExecMode::Unspecified; + //! Submitted script text and syntax. TQueryContent ScriptContent; + //! Script execution statistics. TExecStats ExecStats; + //! Metadata for the script's result sets. std::vector ResultSetsMeta; }; + //! Inherits constructors for operation states that do not contain script metadata. using TOperation::TOperation; + //! Constructs an operation and extracts script metadata from the wire response. TScriptExecutionOperation(TStatus&& status, Ydb::Operations::Operation&& operation); + //! Returns the script execution metadata. const TMetadata& Metadata() const { return Metadata_; } @@ -210,24 +304,36 @@ class TScriptExecutionOperation : public TOperation { TMetadata Metadata_; }; +//! Settings for fetching one page of script results. struct TFetchScriptResultsSettings : public TRequestSettings { + //! Sets the continuation token returned by the previous fetch request. FLUENT_SETTING(std::string, FetchToken); + //! Sets the maximum number of rows to fetch; defaults to 1000. FLUENT_SETTING_DEFAULT(uint64_t, RowsLimit, 1000); + //! Sets a per-request retry policy, overriding the client retry settings when present. FLUENT_SETTING_OPTIONAL(TRetryOperationSettings, RetrySettings); }; +//! Result of fetching one page of script results. class TFetchScriptResultsResult : public TStatus { public: + //! Returns whether this response contains a result set. bool HasResultSet() const { return ResultSet_.has_value(); } + //! Returns the index of the result set in this response. HasResultSet() must be true. uint64_t GetResultSetIndex() const { return ResultSetIndex_; } + //! Returns the result set. HasResultSet() must be true. const TResultSet& GetResultSet() const { return *ResultSet_; } + //! Moves the result set out of this response. HasResultSet() must be true. TResultSet ExtractResultSet() { return std::move(*ResultSet_); } + //! Returns the continuation token for the next page, or an empty string at the end. const std::string& GetNextFetchToken() const { return NextFetchToken_; } + //! Constructs a fetch result that does not contain a result set. explicit TFetchScriptResultsResult(TStatus&& status) : TStatus(std::move(status)) {} + //! Constructs a successful fetch result containing a result set page. TFetchScriptResultsResult(TStatus&& status, TResultSet&& resultSet, int64_t resultSetIndex, const std::string& nextFetchToken) : TStatus(std::move(status)) , ResultSet_(std::move(resultSet)) diff --git a/include/ydb-cpp-sdk/client/query/stats.h b/include/ydb-cpp-sdk/client/query/stats.h index 91b5b54390b..d715e14008b 100644 --- a/include/ydb-cpp-sdk/client/query/stats.h +++ b/include/ydb-cpp-sdk/client/query/stats.h @@ -21,11 +21,15 @@ namespace NYdb::inline V3 { namespace NYdb::inline V3::NQuery { +//! Row and byte counts for one kind of table operation. class TOperationStats { public: + //! Constructs operation statistics from their wire representation. explicit TOperationStats(const Ydb::TableStats::OperationStats& proto); + //! Returns the number of affected rows. uint64_t GetRows() const; + //! Returns the number of affected bytes. uint64_t GetBytes() const; private: @@ -33,15 +37,23 @@ class TOperationStats { uint64_t Bytes_ = 0; }; +//! Statistics for all operations performed on one table. class TTableAccessStats { public: + //! Constructs table access statistics from their wire representation. explicit TTableAccessStats(const Ydb::TableStats::TableAccessStats& proto); + //! Returns the table name. const std::string& GetName() const; + //! Returns read operation statistics. const TOperationStats& GetReads() const; + //! Returns update, insert, upsert, and replace operation statistics. const TOperationStats& GetUpdates() const; + //! Returns delete operation statistics. const TOperationStats& GetDeletes() const; + //! Returns the number of accessed table partitions. uint64_t GetPartitionsCount() const; + //! Returns the number of affected rows when collection was enabled. std::optional GetAffectedRows() const; private: @@ -53,16 +65,25 @@ class TTableAccessStats { std::optional AffectedRows_; }; +//! Statistics for one query execution phase. class TQueryPhaseStats { public: + //! Constructs query phase statistics from their wire representation. explicit TQueryPhaseStats(const Ydb::TableStats::QueryPhaseStats& proto); + //! Returns the phase wall-clock duration in microseconds. uint64_t GetDurationUs() const; + //! Returns the phase wall-clock duration. TDuration GetDuration() const; + //! Returns the phase CPU time in microseconds. uint64_t GetCpuTimeUs() const; + //! Returns the phase CPU time. TDuration GetCpuTime() const; + //! Returns the number of shards affected by the phase. uint64_t GetAffectedShards() const; + //! Returns whether this was a literal execution phase. bool IsLiteralPhase() const; + //! Returns per-table access statistics for this phase. const std::vector& GetTableAccess() const; private: @@ -73,14 +94,21 @@ class TQueryPhaseStats { std::vector TableAccess_; }; +//! Query compilation statistics. class TCompilationStats { public: + //! Constructs compilation statistics from their wire representation. explicit TCompilationStats(const Ydb::TableStats::CompilationStats& proto); + //! Returns whether the compiled query was taken from cache. bool IsFromCache() const; + //! Returns the compilation wall-clock duration in microseconds. uint64_t GetDurationUs() const; + //! Returns the compilation wall-clock duration. TDuration GetDuration() const; + //! Returns compilation CPU time in microseconds. uint64_t GetCpuTimeUs() const; + //! Returns compilation CPU time. TDuration GetCpuTime() const; private: @@ -89,29 +117,46 @@ class TCompilationStats { uint64_t CpuTimeUs_ = 0; }; +//! Execution statistics returned for a query or script. class TExecStats { friend class NYdb::TProtoAccessor; public: + //! Constructs an uninitialized statistics holder used for response metadata. + //! Accessors require the holder to be populated by the SDK. TExecStats() = default; + //! Constructs execution statistics by taking ownership of their wire representation. explicit TExecStats(Ydb::TableStats::QueryStats&& proto); + //! Constructs execution statistics by copying their wire representation. explicit TExecStats(const Ydb::TableStats::QueryStats& proto); + //! Returns the protobuf text representation, optionally including plan, AST, and query metadata. std::string ToString(bool withPlan = false) const; + //! Returns CPU time spent by the query process in microseconds. uint64_t GetProcessCpuTimeUs() const; + //! Returns total query duration in microseconds. uint64_t GetTotalDurationUs() const; + //! Returns total query CPU time in microseconds. uint64_t GetTotalCpuTimeUs() const; + //! Returns the query plan when it was collected. std::optional GetPlan() const; + //! Returns the query abstract syntax tree when it was collected. std::optional GetAst() const; + //! Returns additional query compilation metadata when it was collected. std::optional GetMeta() const; + //! Returns CPU time spent by the query process. TDuration GetProcessCpuTime() const; + //! Returns total query duration. TDuration GetTotalDuration() const; + //! Returns total query CPU time. TDuration GetTotalCpuTime() const; + //! Returns statistics for every query execution phase. std::vector GetQueryPhases() const; + //! Returns compilation statistics when they were collected. std::optional GetCompilation() const; private: diff --git a/include/ydb-cpp-sdk/client/query/tx.h b/include/ydb-cpp-sdk/client/query/tx.h index 843d4ecd0bc..4e7a99aa724 100644 --- a/include/ydb-cpp-sdk/client/query/tx.h +++ b/include/ydb-cpp-sdk/client/query/tx.h @@ -6,48 +6,61 @@ namespace NYdb::inline V3::NQuery { +//! Additional settings for an online read-only transaction. struct TTxOnlineSettings { using TSelf = TTxOnlineSettings; + //! Allows an individual read to observe inconsistent data; disabled by default. FLUENT_SETTING_DEFAULT(bool, AllowInconsistentReads, false); + //! Constructs online read-only settings with consistent individual reads. TTxOnlineSettings() {} }; +//! Selects the isolation and access mode of a query transaction. struct TTxSettings { using TSelf = TTxSettings; + //! Constructs serializable read-write transaction settings. TTxSettings() : Mode_(TS_SERIALIZABLE_RW) {} + //! Creates serializable read-write transaction settings. static TTxSettings SerializableRW() { return TTxSettings(TS_SERIALIZABLE_RW); } + //! Creates online read-only transaction settings. static TTxSettings OnlineRO(const TTxOnlineSettings& settings = TTxOnlineSettings()) { return TTxSettings(TS_ONLINE_RO).OnlineSettings(settings); } + //! Creates stale read-only transaction settings. static TTxSettings StaleRO() { return TTxSettings(TS_STALE_RO); } + //! Creates snapshot read-only transaction settings. static TTxSettings SnapshotRO() { return TTxSettings(TS_SNAPSHOT_RO); } + //! Creates snapshot read-write transaction settings. static TTxSettings SnapshotRW() { return TTxSettings(TS_SNAPSHOT_RW); } + //! Creates read-committed read-write transaction settings. static TTxSettings ReadCommittedRW() { return TTxSettings(TS_READ_COMMITTED_RW); } + //! Creates strict-serializable read-write transaction settings. static TTxSettings StrictSerializableRW() { return TTxSettings(TS_STRICT_SERIALIZABLE_RW); } + //! Writes a human-readable transaction mode name to out. void Out(IOutputStream& out) const { switch (Mode_) { case TS_SERIALIZABLE_RW: @@ -77,18 +90,28 @@ struct TTxSettings { } } + //! Transaction isolation and access modes supported by Query Service. enum ETransactionMode { + //! Serializable read-write mode. TS_SERIALIZABLE_RW, + //! Online read-only mode. TS_ONLINE_RO, + //! Stale read-only mode. TS_STALE_RO, + //! Snapshot read-only mode. TS_SNAPSHOT_RO, + //! Snapshot read-write mode. TS_SNAPSHOT_RW, + //! Read-committed read-write mode. TS_READ_COMMITTED_RW, + //! Strict-serializable read-write mode. TS_STRICT_SERIALIZABLE_RW, }; + //! Sets options used by online read-only mode. FLUENT_SETTING(TTxOnlineSettings, OnlineSettings); + //! Returns the selected transaction mode. ETransactionMode GetMode() const { return Mode_; } diff --git a/include/ydb-cpp-sdk/client/table/table.h b/include/ydb-cpp-sdk/client/table/table.h index 72a062715fa..16350bea3c6 100644 --- a/include/ydb-cpp-sdk/client/table/table.h +++ b/include/ydb-cpp-sdk/client/table/table.h @@ -339,6 +339,8 @@ struct TVectorIndexSettings { Uint8, Int8, Bit, + Float16, + BFloat16, }; EMetric Metric = EMetric::Unspecified; @@ -369,6 +371,8 @@ struct TKMeansTreeSettings { Uint8, Int8, Bit, + Float16, + BFloat16, }; TVectorIndexSettings Settings; @@ -800,12 +804,15 @@ class TTtlDeleteAction {}; class TTtlEvictToExternalStorageAction { public: TTtlEvictToExternalStorageAction(const std::string& storageName); + TTtlEvictToExternalStorageAction(const std::string& storageName, const std::optional& objectKeyPrefix); void SerializeTo(Ydb::Table::EvictionToExternalStorageSettings& proto) const; std::string GetStorage() const; + const std::optional& GetObjectKeyPrefix() const; private: std::string Storage_; + std::optional ObjectKeyPrefix_; }; class TTtlTierSettings { diff --git a/include/ydb-cpp-sdk/client/topic/client.h b/include/ydb-cpp-sdk/client/topic/client.h index eb6ddb0c673..5325b28f3b6 100644 --- a/include/ydb-cpp-sdk/client/topic/client.h +++ b/include/ydb-cpp-sdk/client/topic/client.h @@ -66,6 +66,12 @@ class TTopicClient { TAsyncStatus CommitOffset(const std::string& path, uint64_t partitionId, const std::string& consumerName, uint64_t offset, const TCommitOffsetSettings& settings = {}); + // Reset committed offsets for a consumer on all topic partitions + // (including inactive). Not atomic across partitions; drops any active + // read session for this consumer. + TAsyncStatus ResetOffset(const std::string& path, const std::string& consumerName, + const TResetOffsetSettings& settings); + protected: void OverrideCodec(ECodec codecId, std::unique_ptr&& codecImpl); diff --git a/include/ydb-cpp-sdk/client/topic/control_plane.h b/include/ydb-cpp-sdk/client/topic/control_plane.h index f7f1c77293c..1af0f6ac980 100644 --- a/include/ydb-cpp-sdk/client/topic/control_plane.h +++ b/include/ydb-cpp-sdk/client/topic/control_plane.h @@ -1080,4 +1080,34 @@ struct TCommitOffsetSettings : public TOperationRequestSettings { + TResetOffsetSettings& Earliest() { + Position_ = EPosition::Earliest; + return *this; + } + + TResetOffsetSettings& Latest() { + Position_ = EPosition::Latest; + return *this; + } + + TResetOffsetSettings& FromWrittenAt(TInstant writtenAt) { + Position_ = EPosition::FromWrittenAt; + FromWrittenAt_ = writtenAt; + return *this; + } + + enum class EPosition { + Unspecified, + Earliest, + Latest, + FromWrittenAt, + }; + + EPosition Position_ = EPosition::Unspecified; + TInstant FromWrittenAt_ = TInstant::Zero(); +}; + } // namespace NYdb::NTopic diff --git a/include/ydb-cpp-sdk/client/types/credentials/oidc/credentials.h b/include/ydb-cpp-sdk/client/types/credentials/oidc/credentials.h new file mode 100644 index 00000000000..c111fc2d0dd --- /dev/null +++ b/include/ydb-cpp-sdk/client/types/credentials/oidc/credentials.h @@ -0,0 +1,108 @@ +#pragma once + +#include +#include + +#include + +#include +#include +#include +#include +#include + +namespace NYdb::inline V3::NOidc { + +struct TOAuthToken { + std::string Token; + // An absent expiry means the lifetime is unknown; automatic refresh cannot be scheduled. + std::optional ExpiresAt; + + bool IsValid(TInstant now) const; +}; + +struct TTokenCache { + TOAuthToken AccessToken; + std::optional RefreshToken; +}; + +class ITokenCacher { +public: + virtual ~ITokenCacher(); + // Providers may share a cacher across worker threads. Read() and Write() + // must be thread-safe; interprocess synchronization is not required. + // Calls are synchronous on the provider worker and must return promptly. + // Read failures are treated as cache misses; Write failures leave the token + // usable in memory. Implementations should report persistence errors themselves. + virtual std::optional Read() const = 0; + virtual void Write(const TTokenCache& cache) = 0; +}; + +struct TDeviceAuthInfo { + std::string UserCode; + std::string VerificationUrl; + std::optional VerificationUrlComplete; + TInstant ExpiresAt; +}; + +class IAuthAcceptor { +public: + virtual ~IAuthAcceptor(); + // Called synchronously on a provider worker. Return promptly so polling and + // provider destruction can proceed; do not wait for the user to finish sign-in. + // Copy info before handing it to another thread. A shared acceptor may receive + // concurrent calls from different providers. Exceptions fail authentication. + virtual void Accept(const TDeviceAuthInfo& info) = 0; +}; + +struct TStaticOidcConfig { + std::string AccessToken; + std::optional ExpiresAt; +}; + +struct TClientOidcConfig { + std::string ClientId; + std::string ClientSecret; + // The provider adds "openid" if it is not listed, including for client credentials. + std::vector Scopes; +}; + +struct TDeviceOidcConfig { + // Expiry or denial ends the current sign-in attempt. To try again after user + // interaction, create a new provider (a new factory for parameterless CreateProvider()). + std::string ClientId; + // The provider adds "openid" if it is not listed. + std::vector Scopes; +}; + +using TFlowConfig = std::variant; + +struct TOidcConfig { + using TSelf = TOidcConfig; + + // Must exactly match the issuer advertised by OpenID Discovery, including a trailing slash. + std::string Issuer; + TFlowConfig FlowConfig; + + FLUENT_SETTING(std::shared_ptr, Cacher); + FLUENT_SETTING(std::shared_ptr, Acceptor); +}; + +// Factory identity is stable for the same credentials and custom hook instances. +// Different hook instances isolate independent user sessions. +// Parameterless CreateProvider() reuses one provider; CreateProvider(facility) creates an independent provider for each call. +// Each device provider can prompt if no usable cached credentials exist. Sharing a +// cacher reuses stored tokens but does not coalesce concurrent authorization flows. +// +// GetAuthInfo() blocks until credentials or an error are available; prefer +// GetAuthInfoAsync() when waiting for interactive sign-in. +// +// HTTP runs synchronously on the provider worker. Destruction stops polling and +// joins that worker; active HTTP may wait for socket/connect timeouts (5 s / 30 s). +// DNS resolution is subject to the system resolver's timeout. Transport timeouts +// bound individual socket operations, not the total duration of a streaming response. +// Keep provider/factory owners alive until their hooks and future callbacks return; +// synchronous destruction from those callbacks is not supported. +std::shared_ptr CreateOidcProviderFactory(const TOidcConfig& config); + +} // namespace NYdb::inline V3::NOidc diff --git a/include/ydb-cpp-sdk/client/value/embedding.h b/include/ydb-cpp-sdk/client/value/embedding.h new file mode 100644 index 00000000000..a39fbc89147 --- /dev/null +++ b/include/ydb-cpp-sdk/client/value/embedding.h @@ -0,0 +1,44 @@ +#pragma once + +#include "value.h" + +#include +#include +#include +#include +#include +#include + +namespace NYdb::inline V3 { +namespace NValueHelpers { + +namespace NPrivate { + +template +concept TEmbeddingNumber = (std::integral && sizeof(T) > 1) || std::same_as || std::same_as; + +} // namespace NPrivate + +//! Builds a Bytes value in YDB FloatVector format. Elements are converted to Float32. +//! An empty range produces a single format byte. Declare the query parameter as Bytes. +template + requires std::ranges::sized_range && NPrivate::TEmbeddingNumber> +TValue Embedding(const TRange& values) { + static_assert(sizeof(float) == sizeof(std::uint32_t)); + static_assert(std::numeric_limits::is_iec559); + + std::string bytes; + bytes.reserve(std::ranges::size(values) * sizeof(float) + 1); + for (auto value : values) { + const std::uint32_t bits = std::bit_cast(static_cast(value)); + for (unsigned shift = 0; shift < 32; shift += 8) { + bytes.push_back(static_cast(bits >> shift)); + } + } + bytes.push_back('\x01'); + + return TValueBuilder().Bytes(bytes).Build(); +} + +} // namespace NValueHelpers +} // namespace NYdb::inline V3 diff --git a/library/cpp/CMakeLists.txt b/library/cpp/CMakeLists.txt index b1f5ee76025..da0e2008f9e 100644 --- a/library/cpp/CMakeLists.txt +++ b/library/cpp/CMakeLists.txt @@ -36,6 +36,7 @@ add_subdirectory(monlib/exception) add_subdirectory(monlib/metrics) add_subdirectory(monlib/service) add_subdirectory(openssl/holders) +add_subdirectory(openssl/crypto) add_subdirectory(openssl/init) add_subdirectory(openssl/io) add_subdirectory(openssl/method) diff --git a/library/cpp/containers/cow_string/subst.h b/library/cpp/containers/cow_string/subst.h index 6090ba54b25..a4fb1f1eb07 100644 --- a/library/cpp/containers/cow_string/subst.h +++ b/library/cpp/containers/cow_string/subst.h @@ -4,28 +4,32 @@ #include -/* Replace all occurences of substring `what` with string `with` starting from position `from`. +/** Replace all occurences of substring \p what with string \p with starting from position \p from. * - * @param text String to modify. - * @param what Substring to replace. - * @param with Substring to use as replacement. - * @param from Position at with to start replacement. + * @param[inout] text String to modify. + * @param[in] what Substring to replace. + * @param[in] with Substring to use as replacement. + * @param[in] from Position at with to start replacement. * - * @return Number of replacements occured. + * @return Number of replacements occured. */ +/**@{*/ size_t SubstGlobal(TCowString& text, TStringBuf what, TStringBuf with, size_t from = 0); size_t SubstGlobal(TUtf16CowString& text, TWtringBuf what, TWtringBuf with, size_t from = 0); size_t SubstGlobal(TUtf32CowString& text, TUtf32StringBuf what, TUtf32StringBuf with, size_t from = 0); +/**@}*/ -/* Replace all occurences of character `what` with character `with` starting from position `from`. +/** Replace all occurences of substring \p what with string \p with starting from position \p from. * - * @param text String to modify. - * @param what Character to replace. - * @param with Character to use as replacement. - * @param from Position at with to start replacement. + * @param[inout] text String to modify. + * @param[in] what Character to replace. + * @param[in] with Character to use as replacement. + * @param[in] from Position at with to start replacement. * - * @return Number of replacements occured. + * @return Number of replacements occured. */ +/**@{*/ size_t SubstGlobal(TCowString& text, char what, char with, size_t from = 0); size_t SubstGlobal(TUtf16CowString& text, wchar16 what, wchar16 with, size_t from = 0); size_t SubstGlobal(TUtf32CowString& text, wchar32 what, wchar32 with, size_t from = 0); +/**@}*/ diff --git a/library/cpp/containers/paged_vector/paged_vector.h b/library/cpp/containers/paged_vector/paged_vector.h index 43073852ab3..3b24d4d139a 100644 --- a/library/cpp/containers/paged_vector/paged_vector.h +++ b/library/cpp/containers/paged_vector/paged_vector.h @@ -1,11 +1,11 @@ #pragma once -#include #include #include #include #include +#include namespace NPagedVector { template @@ -172,7 +172,7 @@ namespace NPagedVector { } }; - using TPages = TVector>; + using TPages = TVector>; using TSelf = TPagedVector; TPages Pages_; @@ -199,7 +199,7 @@ namespace NPagedVector { Pages_.reserve(other.Pages_.size()); try { for (auto& ptr : other.Pages_) { - auto& newPage = *Pages_.emplace_back(MakeHolder()); + auto& newPage = *Pages_.emplace_back(std::make_unique()); CurrentPageSize_ = 0; const size_t copyCount = Pages_.size() == other.Pages_.size() ? other.CurrentPageSize_ @@ -345,7 +345,7 @@ namespace NPagedVector { } void AllocateNewPage() { - Pages_.emplace_back(MakeHolder()); + Pages_.emplace_back(std::make_unique()); CurrentPageSize_ = 0; } diff --git a/library/cpp/coroutine/engine/poller.cpp b/library/cpp/coroutine/engine/poller.cpp index 4669828a07e..410cd30ead5 100644 --- a/library/cpp/coroutine/engine/poller.cpp +++ b/library/cpp/coroutine/engine/poller.cpp @@ -173,7 +173,7 @@ namespace { } void Erase(size_t i) noexcept { - V_.Get(i).Destroy(); + V_.Get(i).reset(); } size_t Size() const noexcept { diff --git a/library/cpp/coroutine/listener/listen.cpp b/library/cpp/coroutine/listener/listen.cpp index 3d4e711d1d5..b73162c7033 100644 --- a/library/cpp/coroutine/listener/listen.cpp +++ b/library/cpp/coroutine/listener/listen.cpp @@ -316,7 +316,7 @@ void TContListener::Bind(const TNetworkAddress& addr) { } void TContListener::Stop() noexcept { - Impl_.Destroy(); + Impl_.reset(); } void TContListener::StopListenAddr(const IRemoteAddr& addr) { diff --git a/library/cpp/cpuid_check/README.md b/library/cpp/cpuid_check/README.md index 9c0e8dfa590..fb385a8940c 100644 --- a/library/cpp/cpuid_check/README.md +++ b/library/cpp/cpuid_check/README.md @@ -1,13 +1,13 @@ -Simple utility to check base target x86 SIMD exensions at startup. +Simple utility to check base target x86 SIMD extensions at startup. Program may be built with some SIMD extension enabled (e.g. `-msse4.2`). `PEERDIR` to this library adds statrup check that machine where the program is running supports SIMD extension the program is built for. Currently supported check are: sse4.2, pclmul, aes, avx, avx2 and fma. **Note:** the library depends on `util`. -**Note:** the library adds stratup code and so if `PEERDIR`-ed from `LIBRARY` will do so for all `PROGRAM`-s that (transitively) use the `LIBRARY`. Don't do this! +**Note:** the library adds startup code and so if `PEERDIR`-ed from `LIBRARY` will do so for all `PROGRAM`-s that (transitively) use the `LIBRARY`. Don't do this! You normally don't need to `PEERDIR` this library at all. Since making sse4 in Arcadia default this library is used implicitly. It is `PEERDIR`-ed from all `PROGRAM`-s and derived modules (e.g. `PY2_PROGRAM`, but not `GO_PROGRAM` or `JAVA_PROGRAM`). -It is also not applied to `PROGRAM`-s where `NO_UTIL()`, `NO_PLATFORM()` or `ALLOCATOR(FAKE)` set to avoid undesired dependencied. To disable this implicit check use `NO_CPU_CHECK()` macro or `-DCPU_CHECK=no` ya make flag. +It is also not applied to `PROGRAM`-s where `NO_UTIL()`, `NO_PLATFORM()` or `ALLOCATOR(FAKE)` set to avoid undesired dependencies. To disable this implicit check use `NO_CPU_CHECK()` macro or `-DCPU_CHECK=no` ya make flag. diff --git a/library/cpp/diff/README.md b/library/cpp/diff/README.md index ff68b10eaef..6d5501e700d 100644 --- a/library/cpp/diff/README.md +++ b/library/cpp/diff/README.md @@ -1 +1 @@ -Note: underlying algorithm `library/cpp/lcs` has complexity of O(r log n) by time and O(r) of additional memory, where r is the number of pairs (i, j) for which S1[i] = S2[j]. When comparing file with itself (or with little modifications) it becomes quadratic on the number of occurences of the most frequent line. +Note: underlying algorithm `library/cpp/lcs` has complexity of O(r log n) by time and O(r) of additional memory, where r is the number of pairs (i, j) for which S1[i] = S2[j]. When comparing file with itself (or with little modifications) it becomes quadratic on the number of occurrences of the most frequent line. diff --git a/library/cpp/getopt/small/completer_command.cpp b/library/cpp/getopt/small/completer_command.cpp index 132fdfc4b02..61a9217d573 100644 --- a/library/cpp/getopt/small/completer_command.cpp +++ b/library/cpp/getopt/small/completer_command.cpp @@ -6,9 +6,36 @@ #include +#include +#include + +#include +#include + #include +#include + namespace NLastGetopt { + namespace { + + enum class EShell { + Bash, + Zsh, + }; + + std::optional ParseShell(TStringBuf value) { + if (value == "bash") { + return EShell::Bash; + } + if (value == "zsh") { + return EShell::Zsh; + } + return std::nullopt; + } + + } // namespace + TString MakeInfo(TStringBuf command, TStringBuf flag) { TString info = ( "This command generates shell script with completion function and prints it to `stdout`, " @@ -87,6 +114,81 @@ namespace NLastGetopt { return NComp::Choice({{"zsh"}, {"bash"}}); } + TString MakeModeInfo(const TCompletionConfig& config) { + TStringBuilder info; + info << "Generate shell completion for `" << config.Command << "`"; + for (const auto& alias : config.CommandAliases) { + info << ", `" << alias << "`"; + } + if (!config.YaToolName.empty()) { + info << ", and `ya tool " << config.YaToolName << "`"; + } + if (config.EnableInstaller) { + info << ". Without `--install`, print the script to stdout."; + info << " With `--install`, write the standalone completion"; + if (!config.YaToolName.empty()) { + info << " and the `ya tool` completion shard"; + } + info << " to the user data directory."; + if (!config.YaToolName.empty()) { + info << " Run `ya completion --install --bash` or `ya completion --install --zsh` once to enable " + "completion for the `ya` command itself."; + } + } else { + info << ". Print the script to stdout."; + } + return info; + } + + TString GenerateCompletion( + const TModChooser* modChooser, + const TCompletionConfig& config, + EShell shell) + { + TStringStream output; + switch (shell) { + case EShell::Bash: + TBashCompletionGenerator(modChooser).Generate(config, output); + break; + case EShell::Zsh: + TZshCompletionGenerator(modChooser).Generate(config, output); + break; + } + return std::move(output).Str(); + } + + TVector GetCompletionPaths(const TCompletionConfig& config, EShell shell) { + auto dataHome = GetEnv("XDG_DATA_HOME"); + if (dataHome.empty()) { + dataHome = (TFsPath(GetHomeDir()) / ".local" / "share").GetPath(); + } + + TFsPath directory; + TString fileName; + switch (shell) { + case EShell::Bash: + directory = TFsPath(dataHome) / "bash-completion" / "completions"; + fileName = config.Command; + break; + case EShell::Zsh: + directory = TFsPath(dataHome) / "zsh" / "site-functions"; + fileName = "_" + config.Command; + break; + } + + TVector paths = {directory / fileName}; + if (!config.YaToolName.empty()) { + paths.push_back(directory / "ya-tool.d" / config.YaToolName); + } + return paths; + } + + void WriteCompletion(const TFsPath& path, TStringBuf completion) { + path.Parent().MkDirs(MODE0755); + TFileOutput output(path.GetPath()); + output.Write(completion.data(), completion.size()); + } + TOpt MakeCompletionOpt(const TOpts* opts, TString command, TString name) { return TOpt() .AddLongName(name) @@ -113,53 +215,112 @@ namespace NLastGetopt { class TCompleterMode: public TMainClassArgs { public: - TCompleterMode(const TModChooser* modChooser, TString command, TString modName) - : Command_(std::move(command)) + TCompleterMode(const TModChooser* modChooser, TCompletionConfig config) + : Config_(std::move(config)) , Modes_(modChooser) - , ModName_(std::move(modName)) { } protected: void RegisterOptions(NLastGetopt::TOpts& opts) override { TMainClassArgs::RegisterOptions(opts); + if (Config_.EnableUserFriendlyUsage) { + opts.EnableUserFriendlyUsage(); + if (Modes_->IsSvnRevisionOptionDisabled()) { + if (auto* svnRevisionOption = opts.FindLongOption("svnrevision")) { + svnRevisionOption->Hidden(); + } + } + } opts.SetTitle("Generate tab completion scripts for zsh or bash"); - opts.AddSection("Description", MakeInfo(Command_, ModName_)); + if (Config_.EnableInstaller) { + auto installHelp = TString("Install completion for direct invocations"); + if (!Config_.YaToolName.empty()) { + installHelp += " and ya tool " + Config_.YaToolName; + } + opts.AddLongOption("install", installHelp) + .NoArgument() + .SetFlag(&Install_); + } + + if (Config_.EnableUserFriendlyUsage) { + opts.AddSection("Description", MakeModeInfo(Config_)); + } else { + opts.AddSection("Description", MakeInfo(Config_.Command, Config_.ModName)); + } + + if (Config_.EnableInstaller) { + TStringBuilder examples; + examples << Config_.Command << " " << Config_.ModName << " bash --install"; + if (!Config_.YaToolName.empty()) { + examples << "\nya tool " << Config_.YaToolName << " " << Config_.ModName << " zsh --install"; + } + opts.SetExamples(examples); + } opts.SetFreeArgsNum(1); - opts.GetFreeArgSpec(0) - .Title("") - .Help("shell syntax for completion script (bash or zsh)") - .CompletionArgHelp("shell syntax for completion script") - .Completer(ShellChoiceCompleter()); + auto& shellArg = opts.GetFreeArgSpec(0); + if (Config_.EnableUserFriendlyUsage) { + shellArg + .Title("SHELL") + .Help("Shell whose completion script should be generated: bash or zsh") + .CompletionArgHelp("shell"); + } else { + shellArg + .Title("") + .Help("shell syntax for completion script (bash or zsh)") + .CompletionArgHelp("shell syntax for completion script"); + } + shellArg.Completer(ShellChoiceCompleter()); } int DoRun(NLastGetopt::TOptsParseResult&& parsedOptions) override { auto arg = parsedOptions.GetFreeArgs()[0]; arg.to_lower(); - if (arg == "bash") { - TBashCompletionGenerator(Modes_).Generate(Command_, Cout); - } else if (arg == "zsh") { - TZshCompletionGenerator(Modes_).Generate(Command_, Cout); - } else { + const auto shell = ParseShell(arg); + if (!shell) { Cerr << "Unknown shell name " << arg.Quote() << Endl; parsedOptions.PrintUsage(); return 1; } + auto completion = GenerateCompletion(Modes_, Config_, *shell); + if (!Install_) { + Cout << completion; + return 0; + } + + auto paths = GetCompletionPaths(Config_, *shell); + for (const auto& path : paths) { + WriteCompletion(path, completion); + } + Cout << "Installed " << arg << " completion:" << Endl; + for (const auto& path : paths) { + Cout << " " << path.GetPath() << Endl; + } + Cout << "Restart the shell to load the updated completion." << Endl; return 0; } private: - TString Command_; + TCompletionConfig Config_; const TModChooser* Modes_; - TString ModName_; + bool Install_ = false; }; THolder MakeCompletionMod(const TModChooser* modChooser, TString command, TString modName) { - return MakeHolder(modChooser, std::move(command), std::move(modName)); + return MakeCompletionMod( + modChooser, + TCompletionConfig{ + .ModName = std::move(modName), + .Command = std::move(command), + }); + } + + THolder MakeCompletionMod(const TModChooser* modChooser, TCompletionConfig config) { + return MakeHolder(modChooser, std::move(config)); } } diff --git a/library/cpp/getopt/small/completer_command.h b/library/cpp/getopt/small/completer_command.h index 974cc4617c1..870a56aef67 100644 --- a/library/cpp/getopt/small/completer_command.h +++ b/library/cpp/getopt/small/completer_command.h @@ -8,4 +8,5 @@ namespace NLastGetopt { /// Create a mode that generates completion. THolder MakeCompletionMod(const TModChooser* modChooser, TString command, TString modName = "completion"); + THolder MakeCompletionMod(const TModChooser* modChooser, TCompletionConfig config); } diff --git a/library/cpp/getopt/small/completion_generator.cpp b/library/cpp/getopt/small/completion_generator.cpp index 5e0e55ed38b..e34071bd43b 100644 --- a/library/cpp/getopt/small/completion_generator.cpp +++ b/library/cpp/getopt/small/completion_generator.cpp @@ -1,11 +1,10 @@ #include "completion_generator.h" +#include "last_getopt_parse_result.h" +#include #include #include -#include - -#include "last_getopt_parse_result.h" using NLastGetopt::NEscaping::Q; using NLastGetopt::NEscaping::QQ; @@ -21,6 +20,121 @@ namespace NLastGetopt { #define L out.Line() #define I auto Y_GENERATE_UNIQUE_ID(indent) = out.Indent() + namespace { + + TString MakeShellIdentifier(TStringBuf value) { + TString identifier(value); + for (auto& character : identifier) { + if (!IsAsciiAlnum(character) && character != '_') { + character = '_'; + } + } + return identifier; + } + + struct TCompletionFunctionNames { + TString Main; + TString YaToolExec; + TString YaToolCompletion; + TString YaToolCompletionVariable; + TString YaToolCommandVariable; + }; + + TCompletionFunctionNames MakeCompletionFunctionNames(TStringBuf command) { + const auto yaToolPrefix = "_" + MakeShellIdentifier(command) + "_ya_tool"; + return { + .Main = "_" + TString(command), + .YaToolExec = yaToolPrefix + "_exec", + .YaToolCompletion = yaToolPrefix + "_completion", + .YaToolCompletionVariable = yaToolPrefix + "_completion_active", + .YaToolCommandVariable = yaToolPrefix + "_command", + }; + } + + bool IsCompletableMode(const TModChooser::TMode& mode) { + return !mode.Hidden && !mode.NoCompletion; + } + + bool IsNamedCompletableMode(const TModChooser::TMode& mode) { + return !mode.Name.empty() && IsCompletableMode(mode); + } + + void GenerateZshModeNormalization(TFormattedOutput& out, const TModChooser& chooser) + { + L << "for (( mode_index = 2; mode_index < CURRENT; ++mode_index )); do"; + { + I; + auto& line = L << "case \"${words[mode_index]}\" in "; + TStringBuf separator; + for (const auto& mode : chooser.GetUnsortedModes()) { + if (!IsNamedCompletableMode(*mode)) { + continue; + } + + line << separator << SS(mode->Name); + separator = "|"; + for (const auto& alias : mode->Aliases) { + line << separator << SS(alias); + } + } + line << ")"; + { + I; + L << "if (( mode_index != 2 )); then"; + { + I; + L << "words=(\"${words[1]}\" \"${words[mode_index]}\" " + "\"${(@)words[2,mode_index-1]}\" \"${(@)words[mode_index+1,-1]}\")"; + } + L << "fi"; + L << "break"; + L << ";;"; + } + L << "esac"; + } + L << "done"; + } + + void GenerateBashModeNormalization(TFormattedOutput& out, const TModChooser& chooser) + { + L << "for (( i=1; i < cword; i++ )); do"; + { + I; + auto& line = L << "case \"${words[i]}\" in "; + TStringBuf separator; + for (const auto& mode : chooser.GetUnsortedModes()) { + if (!IsNamedCompletableMode(*mode)) { + continue; + } + + line << separator << BB(mode->Name); + separator = "|"; + for (const auto& alias : mode->Aliases) { + line << separator << BB(alias); + } + } + line << ")"; + { + I; + L << "mode_found=1"; + L << "if (( i != 1 )); then"; + { + I; + L << "words=(\"${words[0]}\" \"${words[i]}\" \"${words[@]:1:i-1}\" " + "\"${words[@]:i+1}\")"; + L << "prev=\"${words[cword - 1]}\""; + } + L << "fi"; + L << "break"; + L << ";;"; + } + L << "esac"; + } + L << "done"; + } + + } // namespace + TCompletionGenerator::TCompletionGenerator(const TModChooser* modChooser) : Options_(modChooser) { @@ -34,23 +148,54 @@ namespace NLastGetopt { } void TZshCompletionGenerator::Generate(TStringBuf command, IOutputStream& stream) { + Generate(TCompletionConfig{ + .Command = TString(command), + }, stream); + } + + void TZshCompletionGenerator::Generate(const TCompletionConfig& config, IOutputStream& stream) { TFormattedOutput out; - NComp::TCompleterManager manager{command}; + NComp::TCompleterManager manager{config.Command}; + const auto names = MakeCompletionFunctionNames(config.Command); - L << "#compdef " << command; + auto& compdef = L << "#compdef " << config.Command; + for (const auto& alias : config.CommandAliases) { + compdef << " " << alias; + } L; - L << "_" << command << "() {"; + L << names.Main << "() {"; { I; + if (!config.YaToolName.empty()) { + L << "if [[ -n ${" << names.YaToolCompletionVariable << ":-} ]]; then"; + { + I; + L << "words=(" << names.YaToolExec << " \"${(@)words[4,-1]}\")"; + L << "(( CURRENT -= 2 ))"; + } + L << "fi"; + L; + } L << "local state line desc modes context curcontext=\"$curcontext\" ret=1"; + if (config.OptionsBeforeMode) { + const auto* modChooser = std::get_if(&Options_); + Y_ABORT_UNLESS(modChooser); + L << "local mode_index"; + GenerateZshModeNormalization(out, **modChooser); + L; + } L << "local words_orig=(\"${words[@]}\")"; L << "local current_orig=\"$((CURRENT - 1))\""; L << "local prefix_orig=\"$PREFIX\""; L << "local suffix_orig=\"$SUFFIX\""; L; std::visit(TOverloaded{ - [&out, &manager](const TModChooser* modChooser) { - GenerateModesCompletion(out, *modChooser, manager); + [&out, &manager, &config](const TModChooser* modChooser) { + GenerateModesCompletion( + out, + *modChooser, + manager, + config.OptionsBeforeMode ? &*config.OptionsBeforeMode : nullptr); }, [&out, &manager](const TOpts* opts) { GenerateOptsCompletion(out, *opts, manager); @@ -63,27 +208,70 @@ namespace NLastGetopt { L; manager.GenerateZsh(out); - // When the completion file is autoloaded by `compinit` from `$fpath`, - // zsh treats the file content as the body of function `_`. - // On first invocation that body merely (re)defines `_` and - // its helpers, so completion would not actually run until the second - // TAB. Calling the redefined function here makes it work on the very - // first TAB and is also harmless when the script is `source`d. - L << "_" << command << " \"$@\""; + if (!config.YaToolName.empty()) { + L << names.YaToolExec << "() {"; + { + I; + L << "command \"$" << names.YaToolCommandVariable << "\" tool " << SS(config.YaToolName) << " \"$@\""; + } + L << "}"; + L; + L << names.YaToolCompletion << "() {"; + { + I; + L << "local " << names.YaToolCompletionVariable << "=1"; + L << "local " << names.YaToolCommandVariable << "=\"$words[1]\""; + L << "local -a words=(\"${words[@]}\")"; + L << "local CURRENT=\"$CURRENT\""; + L << names.Main << " \"$@\""; + } + L << "}"; + L; + L << "__YA_TOOL_COMPLETION_ENTRY=" << names.YaToolCompletion; + L; + L << "if [[ \"$words[1]\" != \"ya\" || \"$words[2]\" != \"tool\" || \"$words[3]\" != " + << SS(config.YaToolName) << " ]]; then"; + { + I; + L << names.Main << " \"$@\""; + } + L << "fi"; + } else { + // When the completion file is autoloaded by `compinit` from `$fpath`, + // zsh treats the file content as the body of function `_`. + // On first invocation that body merely (re)defines `_` and + // its helpers, so completion would not actually run until the second + // TAB. Calling the redefined function here makes it work on the very + // first TAB and is also harmless when the script is `source`d. + L << names.Main << " \"$@\""; + } out.Print(stream); } - void TZshCompletionGenerator::GenerateModesCompletion(TFormattedOutput& out, const TModChooser& chooser, NComp::TCompleterManager& manager) { + void TZshCompletionGenerator::GenerateModesCompletion( + TFormattedOutput& out, + const TModChooser& chooser, + NComp::TCompleterManager& manager, + const TOpts* optionsBeforeMode) + { auto modes = chooser.GetUnsortedModes(); L << "_arguments -C \\"; - L << " '(- : *)'{-h,--help}'[show help information]' \\"; - if (chooser.GetVersionHandler() != nullptr) { - L << " '(- : *)'{-v,--version}'[display version information]' \\"; - } - if (!chooser.IsSvnRevisionOptionDisabled()) { - L << " '(- : *)--svnrevision[show build information]' \\"; + if (optionsBeforeMode) { + for (const auto& option : optionsBeforeMode->GetOpts()) { + if (!option->Hidden_) { + GenerateOptCompletion(out, *optionsBeforeMode, *option, manager); + } + } + } else { + L << " '(- : *)'{-h,--help}'[show help information]' \\"; + if (chooser.GetVersionHandler() != nullptr) { + L << " '(- : *)'{-v,--version}'[display version information]' \\"; + } + if (!chooser.IsSvnRevisionOptionDisabled()) { + L << " '(- : *)--svnrevision[show build information]' \\"; + } } L << " '(-v --version -h --help --svnrevision)1: :->modes' \\"; L << " '(-v --version -h --help --svnrevision)*:: :->args' \\"; @@ -103,7 +291,7 @@ namespace NLastGetopt { L << "desc='modes'"; L << "modes=("; for (auto& mode : modes) { - if (mode->Hidden) { + if (!IsCompletableMode(*mode)) { continue; } if (!mode->Name.empty()) { @@ -148,7 +336,7 @@ namespace NLastGetopt { I; for (auto& mode : modes) { - if (mode->Name.empty() || mode->Hidden) { + if (!IsNamedCompletableMode(*mode)) { continue; } @@ -382,30 +570,71 @@ namespace NLastGetopt { } void TBashCompletionGenerator::Generate(TStringBuf command, IOutputStream& stream) { + Generate(TCompletionConfig{ + .Command = TString(command), + }, stream); + } + + void TBashCompletionGenerator::Generate(const TCompletionConfig& config, IOutputStream& stream) { TFormattedOutput out; - NComp::TCompleterManager manager{command}; + NComp::TCompleterManager manager{config.Command}; + const auto names = MakeCompletionFunctionNames(config.Command); - L << "_" << command << "() {"; + L << names.Main << "() {"; { I; L << "COMPREPLY=()"; L; L << "local i args opts items candidates"; + L << "local mode_found option_arg_completion"; L; L << "local cur prev words cword"; L << "_get_comp_words_by_ref -n \"\\\"'><=;|&(:\" cur prev words cword"; + if (!config.YaToolName.empty()) { + L; + L << "if [[ -n ${" << names.YaToolCompletionVariable << ":-} ]]; then"; + { + I; + L << "words=(" << names.YaToolExec << " \"${words[@]:3}\")"; + L << "(( cword -= 2 ))"; + L << "prev=\"${words[cword - 1]}\""; + } + L << "fi"; + } + if (config.OptionsBeforeMode) { + const auto* modChooser = std::get_if(&Options_); + Y_ABORT_UNLESS(modChooser); + L; + GenerateBashModeNormalization(out, **modChooser); + } L; L << "local need_space=\"1\""; L << "local IFS=$' \\t\\n'"; L; - std::visit(TOverloaded{ - [&out, &manager](const TModChooser* modChooser) { - GenerateModesCompletion(out, *modChooser, manager, 1); - }, - [&out, &manager](const TOpts* opts) { - GenerateOptsCompletion(out, *opts, manager, 1); + if (config.OptionsBeforeMode) { + const auto* modChooser = std::get_if(&Options_); + Y_ABORT_UNLESS(modChooser); + L << "if [[ -z $mode_found ]]; then"; + { + I; + GenerateOptionsBeforeModeCompletion(out, *config.OptionsBeforeMode, **modChooser); } - }, Options_); + L << "else"; + { + I; + GenerateModesCompletion(out, **modChooser, manager, 1); + } + L << "fi"; + } else { + std::visit(TOverloaded{ + [&out, &manager](const TModChooser* modChooser) { + GenerateModesCompletion(out, *modChooser, manager, 1); + }, + [&out, &manager](const TOpts* opts) { + GenerateOptsCompletion(out, *opts, manager, 1); + } + }, Options_); + } L; L; L << "__ltrim_colon_completions \"$cur\""; @@ -432,7 +661,31 @@ namespace NLastGetopt { } L << "}"; L; - L << "complete -o nospace -o default -F _" << command << " " << command; + auto& complete = L << "complete -o nospace -o default -F " << names.Main << " " << config.Command; + for (const auto& alias : config.CommandAliases) { + complete << " " << BB(alias); + } + + if (!config.YaToolName.empty()) { + L; + L << names.YaToolExec << "() {"; + { + I; + L << "command \"$" << names.YaToolCommandVariable << "\" tool " << BB(config.YaToolName) << " \"$@\""; + } + L << "}"; + L; + L << names.YaToolCompletion << "() {"; + { + I; + L << "local " << names.YaToolCompletionVariable << "=1"; + L << "local " << names.YaToolCommandVariable << "=\"${COMP_WORDS[0]}\""; + L << names.Main << " \"$@\""; + } + L << "}"; + L; + L << "__YA_TOOL_COMPLETION_ENTRY=" << names.YaToolCompletion; + } out.Print(stream); } @@ -461,7 +714,7 @@ namespace NLastGetopt { auto& line = L << "COMPREPLY+=( $(compgen -W '"; TStringBuf sep = ""; for (auto& mode : modes) { - if (!mode->Hidden && !mode->NoCompletion) { + if (IsNamedCompletableMode(*mode)) { line << sep << B(mode->Name); sep = " "; } @@ -478,7 +731,7 @@ namespace NLastGetopt { I; for (auto& mode : modes) { - if (mode->Name.empty() || mode->Hidden || mode->NoCompletion) { + if (!IsNamedCompletableMode(*mode)) { continue; } @@ -508,6 +761,83 @@ namespace NLastGetopt { L << "fi"; } + void TBashCompletionGenerator::GenerateOptionsBeforeModeCompletion( + TFormattedOutput& out, + const TOpts& opts, + const TModChooser& chooser) + { + L << "option_arg_completion="; + L << "case ${prev} in"; + { + I; + for (const auto& option : opts.GetOpts()) { + if (option->HasArg_ == EHasArg::NO_ARGUMENT || option->IsHidden()) { + continue; + } + + auto& line = L; + TStringBuf separator; + for (char shortName : option->GetShortNames()) { + line << separator << "'-" << B(TStringBuf(&shortName, 1)) << "'"; + separator = "|"; + } + for (const auto& longName : option->GetLongNames()) { + line << separator << "'--" << B(longName) << "'"; + separator = "|"; + } + line << ")"; + { + I; + L << "option_arg_completion=1"; + if (option->Completer_) { + option->Completer_->GenerateBash(out); + } + L << ";;"; + } + } + } + L << "esac"; + L << "if [[ -z $option_arg_completion ]]; then"; + { + I; + L << "if [[ ${cur} == -* ]]; then"; + { + I; + auto& line = L << "COMPREPLY+=( $(compgen -W '"; + TStringBuf separator; + for (const auto& option : opts.GetOpts()) { + if (option->IsHidden()) { + continue; + } + for (char shortName : option->GetShortNames()) { + line << separator << "-" << B(TStringBuf(&shortName, 1)); + separator = " "; + } + for (const auto& longName : option->GetLongNames()) { + line << separator << "--" << B(longName); + separator = " "; + } + } + line << "' -- ${cur}) )"; + } + L << "else"; + { + I; + auto& line = L << "COMPREPLY+=( $(compgen -W '"; + TStringBuf separator; + for (const auto& mode : chooser.GetUnsortedModes()) { + if (IsNamedCompletableMode(*mode)) { + line << separator << B(mode->Name); + separator = " "; + } + } + line << "' -- ${cur}) )"; + } + L << "fi"; + } + L << "fi"; + } + void TBashCompletionGenerator::GenerateOptsCompletion(TFormattedOutput& out, const TOpts& opts, NComp::TCompleterManager&, size_t level) { auto unorderedOpts = opts.GetOpts(); diff --git a/library/cpp/getopt/small/completion_generator.h b/library/cpp/getopt/small/completion_generator.h index 4241bb7d6cc..af6c4c84bcf 100644 --- a/library/cpp/getopt/small/completion_generator.h +++ b/library/cpp/getopt/small/completion_generator.h @@ -28,9 +28,14 @@ namespace NLastGetopt { public: void Generate(TStringBuf command, IOutputStream& stream) override; + void Generate(const TCompletionConfig& config, IOutputStream& stream); private: - static void GenerateModesCompletion(TFormattedOutput& out, const TModChooser& chooser, NComp::TCompleterManager& manager); + static void GenerateModesCompletion( + TFormattedOutput& out, + const TModChooser& chooser, + NComp::TCompleterManager& manager, + const TOpts* optionsBeforeMode = nullptr); static void GenerateOptsCompletion(TFormattedOutput& out, const TOpts& opts, NComp::TCompleterManager& manager); static void GenerateDefaultOptsCompletion(TFormattedOutput& out, NComp::TCompleterManager& manager); static void GenerateOptCompletion(TFormattedOutput& out, const TOpts& opts, const TOpt& opt, NComp::TCompleterManager& manager); @@ -42,10 +47,15 @@ namespace NLastGetopt { public: void Generate(TStringBuf command, IOutputStream& stream) override; + void Generate(const TCompletionConfig& config, IOutputStream& stream); private: static void GenerateModesCompletion(TFormattedOutput& out, const TModChooser& chooser, NComp::TCompleterManager& manager, size_t level); static void GenerateOptsCompletion(TFormattedOutput& out, const TOpts& opts, NComp::TCompleterManager& manager, size_t level); + static void GenerateOptionsBeforeModeCompletion( + TFormattedOutput& out, + const TOpts& opts, + const TModChooser& chooser); static void GenerateDefaultOptsCompletion(TFormattedOutput& out, NComp::TCompleterManager& manager); }; diff --git a/library/cpp/getopt/small/last_getopt_opt.h b/library/cpp/getopt/small/last_getopt_opt.h index 67a937bdb13..e8dbbcdcc0b 100644 --- a/library/cpp/getopt/small/last_getopt_opt.h +++ b/library/cpp/getopt/small/last_getopt_opt.h @@ -549,7 +549,7 @@ namespace NLastGetopt { * * Note: this only works in zsh. * - * @param arg index of free arg + * @param index index of free arg */ TOpt& IfPresentDisableCompletionForFreeArg(size_t index) { DisableCompletionForFreeArg_.push_back(index); diff --git a/library/cpp/getopt/small/last_getopt_opts.cpp b/library/cpp/getopt/small/last_getopt_opts.cpp index b656e607e85..085e498b294 100644 --- a/library/cpp/getopt/small/last_getopt_opts.cpp +++ b/library/cpp/getopt/small/last_getopt_opts.cpp @@ -445,7 +445,9 @@ namespace NLastGetopt { os << "(values: " << choicesHelp << ")"; } - if (opt->HasDefaultValue()) { + if (opt->HasDefaultValue() && + (ShowDefaultValuesForNoArgumentOptions_ || opt->GetHasArg() != NO_ARGUMENT)) + { auto quotedDef = QuoteForHelp(opt->GetDefaultValue()); if (helpHasParagraphs) { os << Endl << Endl << SPad << leftPadding << " "; diff --git a/library/cpp/getopt/small/last_getopt_opts.h b/library/cpp/getopt/small/last_getopt_opts.h index 868477ec79e..7cd66a218ce 100644 --- a/library/cpp/getopt/small/last_getopt_opts.h +++ b/library/cpp/getopt/small/last_getopt_opts.h @@ -69,6 +69,9 @@ namespace NLastGetopt { TString CustomUsage; // user defined usage string TVector> Sections; // additional help entries to print after usage + bool ShowDefaultValuesForNoArgumentOptions_ = true; + bool ShowFreeArgTitlesInErrors_ = false; + bool ShowExceptionTypeInUsageErrors_ = true; public: /** @@ -218,6 +221,15 @@ namespace NLastGetopt { return GetLongOption(name); } + /// @} + + /** + * Search for the option with given short name + * @param c short name for search + * @return ref on result (throw exception if not found) + */ + /// @{ + const TOpt& GetOption(char c) const { return GetCharOption(c); } @@ -431,7 +443,7 @@ namespace NLastGetopt { /** * Replace help string with given * - * @param decr new help string + * @param descr new help string */ void SetCmdLineDescr(const TString& descr) { CustomCmdLineDescr = descr; @@ -446,6 +458,13 @@ namespace NLastGetopt { CustomUsage = usage; } + /** + * Hide default values for options that take no arguments. + */ + void HideDefaultValuesForNoArgumentOptions() { + ShowDefaultValuesForNoArgumentOptions_ = false; + } + /** * Add a section to print after the main usage spec. */ @@ -497,6 +516,29 @@ namespace NLastGetopt { return FreeArgsMin_; } + /** + * Name missing positional arguments in usage errors. + */ + void ShowFreeArgTitlesInErrors() { + ShowFreeArgTitlesInErrors_ = true; + } + + /** + * Hide the exception type prefix in usage error messages. + */ + void HideExceptionTypeInUsageErrors() { + ShowExceptionTypeInUsageErrors_ = false; + } + + /** + * Enable concise help and informative usage errors. + */ + void EnableUserFriendlyUsage() { + HideDefaultValuesForNoArgumentOptions(); + ShowFreeArgTitlesInErrors(); + HideExceptionTypeInUsageErrors(); + } + /** * Set maximal number of free args * diff --git a/library/cpp/getopt/small/last_getopt_parse_result.cpp b/library/cpp/getopt/small/last_getopt_parse_result.cpp index 016fe347141..39f61935540 100644 --- a/library/cpp/getopt/small/last_getopt_parse_result.cpp +++ b/library/cpp/getopt/small/last_getopt_parse_result.cpp @@ -210,7 +210,17 @@ namespace NLastGetopt { } void TOptsParseResult::HandleError() const { - Cerr << CurrentExceptionMessage() << Endl; + TString message; + try { + throw; + } catch (const TUsageException& error) { + message = Parser_.Get() && !Parser_->Opts_->ShowExceptionTypeInUsageErrors_ + ? error.what() + : CurrentExceptionMessage(); + } catch (...) { + message = CurrentExceptionMessage(); + } + Cerr << message << Endl; if (Parser_.Get()) { // parser initializing can fail (and we get here, see Init) if (Parser_->Opts_->FindLongOption("help") != nullptr) { Cerr << "Try '" << Parser_->ProgramName_ << " --help' for more information." << Endl; diff --git a/library/cpp/getopt/small/last_getopt_parser.cpp b/library/cpp/getopt/small/last_getopt_parser.cpp index 98c3951a1c7..501e6f57981 100644 --- a/library/cpp/getopt/small/last_getopt_parser.cpp +++ b/library/cpp/getopt/small/last_getopt_parser.cpp @@ -2,7 +2,10 @@ #include +#include + #include +#include namespace NLastGetopt { void TOptsParser::Init(const TOpts* opts, int argc, const char* argv[]) { @@ -65,11 +68,24 @@ namespace NLastGetopt { Y_ASSERT(!Stopped_); - if (Opts_->FreeArgsMin_ == Opts_->FreeArgsMax_ && Argc_ - Pos_ != Opts_->FreeArgsMin_) + const size_t freeArgCount = Argc_ - Pos_; + if (Opts_->ShowFreeArgTitlesInErrors_ && freeArgCount < Opts_->FreeArgsMin_) { + TVector missingArguments; + for (size_t index = freeArgCount; index < Opts_->FreeArgsMin_; ++index) { + missingArguments.push_back(Opts_->GetFreeArgTitle(index)); + } + throw TUsageException() + << "missing required positional argument" + << (missingArguments.size() == 1 ? "" : "s") + << ": " + << JoinSeq(" ", missingArguments); + } + + if (Opts_->FreeArgsMin_ == Opts_->FreeArgsMax_ && freeArgCount != Opts_->FreeArgsMin_) throw TUsageException() << "required exactly " << Opts_->FreeArgsMin_ << " free args"; - else if (Argc_ - Pos_ < Opts_->FreeArgsMin_) + else if (freeArgCount < Opts_->FreeArgsMin_) throw TUsageException() << "required at least " << Opts_->FreeArgsMin_ << " free args"; - else if (Argc_ - Pos_ > Opts_->FreeArgsMax_) + else if (freeArgCount > Opts_->FreeArgsMax_) throw TUsageException() << "required at most " << Opts_->FreeArgsMax_ << " free args"; return false; diff --git a/library/cpp/getopt/small/modchooser.cpp b/library/cpp/getopt/small/modchooser.cpp index d7fa9dbf511..275d9a07829 100644 --- a/library/cpp/getopt/small/modchooser.cpp +++ b/library/cpp/getopt/small/modchooser.cpp @@ -145,6 +145,23 @@ void TModChooser::SetDescription(const TString& descr) { Description = descr; } +void TModChooser::SetCmdLineDescription(const TString& description) { + CmdLineDescription = description; +} + +void TModChooser::SetModesTitle(const TString& title) { + ModesTitle = title; +} + +void TModChooser::SetModeName(const TString& name, const TString& usageName) { + ModeName = name; + ModeUsageName = usageName; +} + +void TModChooser::SetExamples(const TString& examples) { + Examples = examples; +} + void TModChooser::SetModesHelpOption(const TString& helpOption) { ModesHelpOption = helpOption; } @@ -170,9 +187,33 @@ void TModChooser::DisableSvnRevisionOption() { } void TModChooser::AddCompletions(TString progName, const TString& name, bool hidden, bool noCompletion) { + AddCompletions( + NLastGetopt::TCompletionConfig{ + .Command = std::move(progName), + }, + name, + hidden, + noCompletion); +} + +void TModChooser::AddCompletions( + NLastGetopt::TCompletionConfig config, + const TString& name, + bool hidden, + bool noCompletion) +{ if (CompletionsGenerator == nullptr) { - CompletionsGenerator = NLastGetopt::MakeCompletionMod(this, std::move(progName), name); - AddMode(name, CompletionsGenerator.Get(), "generate autocompletion files", hidden, noCompletion); + TString description; + if (config.EnableInstaller) { + description = "Generate and install shell completion scripts"; + } else if (!config.CommandAliases.empty() || !config.YaToolName.empty()) { + description = "Generate shell completion scripts"; + } else { + description = "generate autocompletion files"; + } + config.ModName = name; + CompletionsGenerator = NLastGetopt::MakeCompletionMod(this, std::move(config)); + AddMode(name, CompletionsGenerator.Get(), description, hidden, noCompletion); } } @@ -311,47 +352,73 @@ TString TModChooser::TMode::FormatFullName(size_t pad, const NColorizer::TColors } void TModChooser::PrintHelp(const TString& progName, bool toStdErr) const { + PrintHelpImpl(progName, toStdErr, false); +} + +void TModChooser::PrintBriefHelp(const TString& progName, bool toStdErr) const { + PrintHelpImpl(progName, toStdErr, true); +} + +void TModChooser::PrintHelpImpl(const TString& progName, bool toStdErr, bool brief) const { auto baseName = TFsPath(progName).Basename(); auto& out = toStdErr ? Cerr : Cout; const auto& colors = toStdErr ? NColorizer::StdErr() : NColorizer::StdOut(); out << Description << Endl << Endl; - out << colors.BoldColor() << "Usage" << colors.OldColor() << ": " << baseName << " MODE [MODE_OPTIONS]" << Endl; + out << colors.BoldColor() << "Usage" << colors.OldColor() << ": " << baseName << " " + << (CmdLineDescription ? CmdLineDescription : "MODE [MODE_OPTIONS]") << Endl; out << Endl; - out << colors.BoldColor() << "Modes" << colors.OldColor() << ":" << Endl; - size_t maxModeLen = 0; - for (const auto& [name, mode] : Modes) { - if (name != mode->Name) - continue; // this is an alias - maxModeLen = Max(maxModeLen, mode->CalculateFullNameLen()); - } + if (brief) { + out << "Run '" << baseName << " --help' for full help." << Endl; + } else { + out << colors.BoldColor() << (ModesTitle ? ModesTitle : "Modes") << colors.OldColor() << ":" << Endl; + size_t maxModeLen = 0; + for (const auto& [name, mode] : Modes) { + if (name != mode->Name) { + continue; // this is an alias + } + maxModeLen = Max(maxModeLen, mode->CalculateFullNameLen()); + } - if (ShowSeparated) { - for (const auto& unsortedMode : UnsortedModes) - if (!unsortedMode->Hidden) { - if (unsortedMode->Name.size()) { - out << " " << unsortedMode->FormatFullName(maxModeLen + 4, colors) << unsortedMode->Description << Endl; - } else { + if (ShowSeparated) { + for (const auto& unsortedMode : UnsortedModes) { + if (unsortedMode->Hidden) { + continue; + } + if (unsortedMode->Name.empty()) { out << SeparationString << Endl; out << unsortedMode->Description << Endl; + continue; } + out << " " << unsortedMode->FormatFullName(maxModeLen + 4, colors) + << unsortedMode->Description << Endl; } - } else { - for (const auto& mode : Modes) { - if (mode.first != mode.second->Name) - continue; // this is an alias - - if (!mode.second->Hidden) { - out << " " << mode.second->FormatFullName(maxModeLen + 4, colors) << mode.second->Description << Endl; + } else { + for (const auto& [name, mode] : Modes) { + if (name == mode->Name && !mode->Hidden) { + out << " " << mode->FormatFullName(maxModeLen + 4, colors) << mode->Description << Endl; + } } } - } - out << Endl; - out << "To get help for specific mode type '" << baseName << " MODE " << ModesHelpOption << "'" << Endl; - if (VersionHandler) - out << "To print program version type '" << baseName << " --version'" << Endl; - if (!SvnRevisionOptionDisabled) { - out << "To print svn revision type '" << baseName << " --svnrevision'" << Endl; + out << Endl; + if (ModeName) { + out << "Run '" << baseName << " " << ModeUsageName << " " << ModesHelpOption + << "' for help with a specific " + << ModeName << "." << Endl; + } else { + out << "To get help for specific mode type '" << baseName << " MODE " << ModesHelpOption << "'" << Endl; + } + if (VersionHandler) { + out << "To print program version type '" << baseName << " --version'" << Endl; + } + if (!SvnRevisionOptionDisabled) { + out << "To print svn revision type '" << baseName << " --svnrevision'" << Endl; + } + } + if (Examples) { + out << Endl; + out << colors.BoldColor() << "Examples" << colors.OldColor() << ":" << Endl; + out << Examples << Endl; } } diff --git a/library/cpp/getopt/small/modchooser.h b/library/cpp/getopt/small/modchooser.h index c41ac3e5608..7135a22e09d 100644 --- a/library/cpp/getopt/small/modchooser.h +++ b/library/cpp/getopt/small/modchooser.h @@ -7,8 +7,24 @@ #include #include +#include #include +namespace NLastGetopt { + +struct TCompletionConfig +{ + TString ModName = "completion"; + TString Command; + TVector CommandAliases; + TString YaToolName; + bool EnableInstaller = false; + bool EnableUserFriendlyUsage = false; + std::optional OptionsBeforeMode; +}; + +} // namespace NLastGetopt + //! Mode function with vector of cli arguments. using TMainFunctionPtrV = std::function&)> ; using TMainFunctionRawPtrV = int (*)(const TVector& argv); @@ -79,6 +95,18 @@ class TModChooser { //! Set main program description. void SetDescription(const TString& descr); + //! Replace the command-line description in the usage block. + void SetCmdLineDescription(const TString& description); + + //! Set the title above the mode list. + void SetModesTitle(const TString& title); + + //! Set the singular mode name and its usage placeholder. + void SetModeName(const TString& name, const TString& usageName); + + //! Set examples shown at the bottom of help output. + void SetExamples(const TString& examples); + //! Set modes help option name (-? is by default) void SetModesHelpOption(const TString& helpOption); @@ -102,6 +130,11 @@ class TModChooser { void DisableSvnRevisionOption(); void AddCompletions(TString progName, const TString& name = "completion", bool hidden = false, bool noCompletion = false); + void AddCompletions( + NLastGetopt::TCompletionConfig config, + const TString& name = "completion", + bool hidden = false, + bool noCompletion = false); void SetSubcommandPath(const TVector& subcommandPath) const; const TVector& GetSubcommandPath() const; @@ -126,6 +159,9 @@ class TModChooser { void PrintHelp(const TString& progName, bool toStdErr = false) const; + //! Print description, usage, and examples without the mode list. + void PrintBriefHelp(const TString& progName, bool toStdErr = false) const; + struct TMode { TString Name; TMainClass* Main; @@ -159,9 +195,26 @@ class TModChooser { bool IsSvnRevisionOptionDisabled() const; private: + void PrintHelpImpl(const TString& progName, bool toStdErr, bool brief) const; + //! Main program description. TString Description; + //! Command-line description shown after the program name. + TString CmdLineDescription; + + //! Title shown above the mode list. + TString ModesTitle; + + //! Singular mode name used in help text. + TString ModeName; + + //! Mode placeholder used in help commands. + TString ModeUsageName; + + //! Examples shown at the bottom of help output. + TString Examples; + //! Help option for modes. TString ModesHelpOption; diff --git a/library/cpp/getopt/small/posix_getopt.cpp b/library/cpp/getopt/small/posix_getopt.cpp index bd06f3499f7..7a73eb528cc 100644 --- a/library/cpp/getopt/small/posix_getopt.cpp +++ b/library/cpp/getopt/small/posix_getopt.cpp @@ -1,8 +1,7 @@ #include "posix_getopt.h" -#include - #include +#include namespace NLastGetopt { char* optarg; @@ -11,8 +10,8 @@ namespace NLastGetopt { int opterr; int optreset; - static THolder Opts; - static THolder OptsParser; + static std::unique_ptr Opts; + static std::unique_ptr OptsParser; int getopt_long_impl(int argc, char* const* argv, const char* optstring, const struct option* longopts, int* longindex, bool long_only) { @@ -21,7 +20,7 @@ namespace NLastGetopt { optind = 1; opterr = 1; optreset = 0; - Opts.Reset(new TOpts(TOpts::Default(optstring))); + Opts.reset(new TOpts(TOpts::Default(optstring))); Opts->AllowSingleDashForLong_ = long_only; @@ -38,7 +37,7 @@ namespace NLastGetopt { opt->UserValue(o->flag); } - OptsParser.Reset(new TOptsParser(&*Opts, argc, (const char**)argv)); + OptsParser.reset(new TOptsParser(&*Opts, argc, (const char**)argv)); } optarg = nullptr; diff --git a/library/cpp/getopt/small/ygetopt.cpp b/library/cpp/getopt/small/ygetopt.cpp index 1f52827f742..8686081f91b 100644 --- a/library/cpp/getopt/small/ygetopt.cpp +++ b/library/cpp/getopt/small/ygetopt.cpp @@ -5,6 +5,8 @@ #include #include +#include + class TGetOpt::TImpl: public TSimpleRefCount { public: inline TImpl(int argc, const char* const* argv, const TString& fmt) @@ -36,7 +38,7 @@ class TGetOpt::TIterator::TIterImpl: public TSimpleRefCount { } ArgsPtrs_.Get()[Args_.size()] = nullptr; - Opt_.Reset(new Opt((int)Args_.size(), ArgsPtrs_.Get(), Format_.data())); + Opt_.reset(new Opt((int)Args_.size(), ArgsPtrs_.Get(), Format_.data())); } inline ~TIterImpl() = default; @@ -62,7 +64,7 @@ class TGetOpt::TIterator::TIterImpl: public TSimpleRefCount { TVector Args_; TArrayHolder ArgsPtrs_; const TString Format_; - THolder Opt_; + std::unique_ptr Opt_; int OptLet_; const char* Arg_; }; diff --git a/library/cpp/http/fetch/sockhandler.h b/library/cpp/http/fetch/sockhandler.h index e18149f6571..01def0e6da8 100644 --- a/library/cpp/http/fetch/sockhandler.h +++ b/library/cpp/http/fetch/sockhandler.h @@ -97,7 +97,7 @@ class TSimpleSocketHandler { if (!Socket) return; Socket->ShutDown(SHUT_RDWR); - Socket.Destroy(); + Socket.reset(); } void SetSocket(SOCKET fd) { diff --git a/library/cpp/http/io/chunk.cpp b/library/cpp/http/io/chunk.cpp index 6975d9eac1e..42cbbe24ea6 100644 --- a/library/cpp/http/io/chunk.cpp +++ b/library/cpp/http/io/chunk.cpp @@ -241,6 +241,6 @@ void TChunkedOutput::DoFlush() { void TChunkedOutput::DoFinish() { if (Impl_.Get()) { Impl_->Finish(); - Impl_.Destroy(); + Impl_.reset(); } } diff --git a/library/cpp/http/io/headers.cpp b/library/cpp/http/io/headers.cpp index 9117bc0ba00..80cbd9dfa45 100644 --- a/library/cpp/http/io/headers.cpp +++ b/library/cpp/http/io/headers.cpp @@ -42,6 +42,11 @@ void THttpInputHeader::OutTo(IOutputStream* stream) const { stream->Write(parts, sizeof(parts) / sizeof(*parts)); } +template <> +void Out(IOutputStream& out, const THttpInputHeader& h) { + h.OutTo(&out); +} + THttpHeaders::THttpHeaders(IInputStream* stream) { TString header; TString line; diff --git a/library/cpp/http/io/headers.h b/library/cpp/http/io/headers.h index 72d3dec4bb0..343f253d3eb 100644 --- a/library/cpp/http/io/headers.h +++ b/library/cpp/http/io/headers.h @@ -6,6 +6,8 @@ #include #include +#include + class IInputStream; class IOutputStream; @@ -136,4 +138,6 @@ class THttpHeaders { THeaders Headers_; }; +using TEncodeContentPredicate = std::function; + /// @} diff --git a/library/cpp/http/io/stream.cpp b/library/cpp/http/io/stream.cpp index a97721c722a..4cda41d021f 100644 --- a/library/cpp/http/io/stream.cpp +++ b/library/cpp/http/io/stream.cpp @@ -620,6 +620,10 @@ class THttpOutput::TImpl { CompressionHeaderEnabled_ = enable; } + inline void SetContentEncodingPredicate(TEncodeContentPredicate predicate) { + ContentEncodingPredicate_ = std::move(predicate); + } + inline bool IsCompressionEnabled() const noexcept { return !ComprSchemas_.empty(); } @@ -793,7 +797,8 @@ class THttpOutput::TImpl { } if (IsHttpResponse()) { - if (Request_ && IsCompressionEnabled() && HasResponseBody()) { + const bool contentEncodingAllowed = IsContentEncodingAllowed(); + if (Request_ && IsCompressionEnabled() && HasResponseBody() && contentEncodingAllowed) { TString scheme = Request_->BestCompressionScheme(ComprSchemas_); if (scheme != "identity") { AddOrReplaceHeader(THttpInputHeader("Content-Encoding", scheme)); @@ -801,7 +806,7 @@ class THttpOutput::TImpl { } } - RebuildStream(); + RebuildStream(contentEncodingAllowed); } else { if (IsCompressionEnabled()) { AddOrReplaceHeader(THttpInputHeader("Accept-Encoding", BuildAcceptEncoding())); @@ -826,7 +831,16 @@ class THttpOutput::TImpl { return ret; } - inline void RebuildStream() { + inline bool IsContentEncodingAllowed() const { + if (!ContentEncodingPredicate_) { + return true; + } + + static const THttpHeaders emptyHeaders; + return ContentEncodingPredicate_(Request_ ? Request_->Headers() : emptyHeaders, Headers_); + } + + inline void RebuildStream(bool contentEncodingAllowed = true) { bool keepAlive = false; const TCompressionCodecFactory::TEncoderConstructor* encoder = nullptr; bool chunked = false; @@ -838,7 +852,10 @@ class THttpOutput::TImpl { if (hl == TStringBuf("connection")) { keepAlive = to_lower(header.Value()) == TStringBuf("keep-alive"); - } else if (IsCompressionHeaderEnabled() && hl == TStringBuf("content-encoding")) { + } else if (contentEncodingAllowed && + IsCompressionHeaderEnabled() && + hl == TStringBuf("content-encoding")) + { encoder = TCompressionCodecFactory::Instance().FindEncoder(to_lower(header.Value())); } else if (hl == TStringBuf("transfer-encoding")) { chunked = to_lower(header.Value()) == TStringBuf("chunked"); @@ -887,6 +904,7 @@ class THttpOutput::TImpl { size_t Version_; TArrayRef ComprSchemas_; + TEncodeContentPredicate ContentEncodingPredicate_; bool KeepAliveEnabled_; bool BodyEncodingEnabled_; @@ -955,6 +973,10 @@ void THttpOutput::EnableCompressionHeader(bool enable) { Impl_->EnableCompressionHeader(enable); } +void THttpOutput::SetContentEncodingPredicate(TEncodeContentPredicate predicate) { + Impl_->SetContentEncodingPredicate(std::move(predicate)); +} + bool THttpOutput::IsKeepAliveEnabled() const noexcept { return Impl_->IsKeepAliveEnabled(); } diff --git a/library/cpp/http/io/stream.h b/library/cpp/http/io/stream.h index 5ea5e18270e..1871f6e751d 100644 --- a/library/cpp/http/io/stream.h +++ b/library/cpp/http/io/stream.h @@ -143,6 +143,10 @@ class THttpOutput: public IOutputStream { /// указанным в Content-Encoding (включен по умолчанию) void EnableCompressionHeader(bool enable); + /// Устанавливает политику HTTP-кодирования тела ответа. Если предикат вернул false, + /// сжатие тела не применяется, а заголовок Content-Encoding передается без изменений. + void SetContentEncodingPredicate(TEncodeContentPredicate predicate); + /// Проверяет, производится ли выдача ответов в упакованном виде. bool IsCompressionEnabled() const noexcept; diff --git a/library/cpp/http/io/stream_ut.cpp b/library/cpp/http/io/stream_ut.cpp index 9eb59acab74..9c6cab72dda 100644 --- a/library/cpp/http/io/stream_ut.cpp +++ b/library/cpp/http/io/stream_ut.cpp @@ -441,6 +441,67 @@ Y_UNIT_TEST_SUITE(THttpStreamTest) { UNIT_ASSERT(!result.Contains("content-length")); } + TString MakeCompressedResponse(bool allowContentEncoding, TStringBuf contentEncoding = {}) { + TString requestData = "GET / HTTP/1.1\r\nAccept-Encoding: gzip\r\n"; + requestData += "X-Allow-Content-Encoding: "; + requestData += allowContentEncoding ? "yes" : "no"; + requestData += "\r\n"; + requestData += "\r\n"; + + TMemoryInput request(requestData); + THttpInput httpInput(&request); + TString result; + TStringOutput output(result); + THttpOutput httpOutput(&output, &httpInput); + httpOutput.EnableCompression(true); + httpOutput.SetContentEncodingPredicate([](const THttpHeaders& requestHeaders, const THttpHeaders& responseHeaders) { + const auto* requestHeader = requestHeaders.FindHeader("X-Allow-Content-Encoding"); + const auto* responseHeader = responseHeaders.FindHeader("X-Allow-Content-Encoding"); + return requestHeader && requestHeader->Value() == "yes" && + responseHeader && responseHeader->Value() == "yes"; + }); + + constexpr TStringBuf body = "Mary had a little lamb."; + httpOutput << "HTTP/1.1 200 OK\r\n" + << "Content-Type: application/octet-stream\r\n" + << "X-Allow-Content-Encoding: " << (allowContentEncoding ? "yes" : "no") << "\r\n" + << "Content-Length: " << body.size() << "\r\n"; + if (contentEncoding) { + httpOutput << "Content-Encoding: " << contentEncoding << "\r\n"; + } + httpOutput << "\r\n" << body; + httpOutput.Finish(); + return result; + } + + void AssertCompressionBypassed(const TString& response) { + TMemoryInput input(response); + THttpInput httpInput(&input); + UNIT_ASSERT(!httpInput.Headers().HasHeader("Content-Encoding")); + UNIT_ASSERT(httpInput.Headers().HasHeader("Content-Length")); + UNIT_ASSERT_VALUES_EQUAL(httpInput.ReadAll(), "Mary had a little lamb."); + } + + Y_UNIT_TEST(ContentEncodingPredicateBypassesNegotiatedCompression) { + AssertCompressionBypassed(MakeCompressedResponse(false)); + } + + Y_UNIT_TEST(ContentEncodingPredicateBypassesExplicitContentEncoding) { + const TString response = MakeCompressedResponse(false, "gzip"); + const TString lower = to_lower(response); + UNIT_ASSERT(lower.Contains("content-encoding: gzip")); + UNIT_ASSERT(lower.Contains("content-length:")); + UNIT_ASSERT(response.EndsWith("Mary had a little lamb.")); + } + + Y_UNIT_TEST(ContentEncodingPredicateAllowsCompression) { + const TString response = MakeCompressedResponse(true); + const TString lower = to_lower(response); + UNIT_ASSERT(lower.Contains("content-encoding: gzip")); + UNIT_ASSERT(!lower.Contains("content-length:")); + UNIT_ASSERT(!response.EndsWith("Mary had a little lamb.")); + } + Y_UNIT_TEST(CodecsPriority) { TMemoryInput request("GET / HTTP/1.1\r\nAccept-Encoding: gzip, br\r\n\r\n"); TVector codecs = {"br", "gzip"}; diff --git a/library/cpp/http/server/http.cpp b/library/cpp/http/server/http.cpp index 79a1194c770..de092955307 100644 --- a/library/cpp/http/server/http.cpp +++ b/library/cpp/http/server/http.cpp @@ -24,6 +24,21 @@ using namespace NAddr; namespace { + class THttpAcceptError: public yexception { + public: + explicit THttpAcceptError(int errorCode) + : ErrorCode_(errorCode) + { + } + + int ErrorCode() const noexcept { + return ErrorCode_; + } + + private: + int ErrorCode_; + }; + class IPollAble { public: inline IPollAble() noexcept { @@ -278,8 +293,8 @@ class THttpServer::TImpl { Connections->Clear(); } - Connections.Destroy(); - Poller.Destroy(); + Connections.reset(); + Poller.reset(); } void Shutdown() { @@ -347,13 +362,14 @@ class THttpServer::TImpl { void OnPollEvent(TInstant) override { SOCKET s = ::accept(S_, nullptr, nullptr); + const int errorCode = s == INVALID_SOCKET ? WSAGetLastError() : 0; if (Server_->Options_.OneShotPoll) { Server_->Poller->WaitReadOneShot(S_, this); } if (s == INVALID_SOCKET) { - ythrow yexception() << "accept: " << LastSystemErrorText(); + ythrow THttpAcceptError(errorCode) << "accept: " << LastSystemErrorText(errorCode); } Server_->AddRequestFromSocket(s, TInstant::Now(), SockAddrRef_); @@ -403,6 +419,8 @@ class THttpServer::TImpl { } } catch (const TShouldStop&) { break; + } catch (const THttpAcceptError& error) { + Cb_->OnAcceptException(error.ErrorCode()); } catch (...) { Cb_->OnException(); } @@ -721,7 +739,7 @@ void TClientRequest::ResetConnection() { if (HttpConn_) { // send RST packet to client HttpConn_->Reset(); - HttpConn_.Destroy(); + HttpConn_.reset(); } } @@ -742,6 +760,7 @@ void TClientRequest::Process(void* ThreadSpecificResource) { auto maxRequestsPerConnection = HttpServ()->Options().MaxRequestsPerConnection; HttpConn_->Output()->EnableKeepAlive(HttpServ()->Options().KeepAliveEnabled && (!maxRequestsPerConnection || Conn_->ReceivedRequests < maxRequestsPerConnection)); HttpConn_->Output()->EnableCompression(HttpServ()->Options().CompressionEnabled); + HttpConn_->Output()->SetContentEncodingPredicate(HttpServ()->Options().ContentEncodingPredicate); } if (!BeforeParseRequestOk(ThreadSpecificResource)) { diff --git a/library/cpp/http/server/http.h b/library/cpp/http/server/http.h index 5b23cb21a00..117a95822cc 100644 --- a/library/cpp/http/server/http.h +++ b/library/cpp/http/server/http.h @@ -39,6 +39,10 @@ class THttpServer { virtual void OnException() { } + virtual void OnAcceptException(int /*errorCode*/) { + OnException(); + } + virtual void OnMaxConn() { } diff --git a/library/cpp/http/server/http_ut.cpp b/library/cpp/http/server/http_ut.cpp index b3280745929..c5069b97b77 100644 --- a/library/cpp/http/server/http_ut.cpp +++ b/library/cpp/http/server/http_ut.cpp @@ -5,6 +5,10 @@ #include #include +#include +#include + +#include #include #include #include @@ -1125,4 +1129,104 @@ Y_UNIT_TEST_SUITE(THttpServerTest) { UNIT_ASSERT_STRINGS_EQUAL(server.InputCopy.Str(), TStringBuilder() << "GET / HTTP/1.1\r\nHost: localhost:" << port << "\r\nConnection: Keep-Alive\r\n\r\n"); UNIT_ASSERT_STRINGS_EQUAL(server.OutputCopy.Str(), TStringBuilder() << "HTTP/1.1 200 Ok\r\nConnection: Keep-Alive\r\nTransfer-Encoding: chunked\r\n\r\n0\r\n\r\n"); } + +#ifdef _linux_ + // The public server API deliberately does not expose listening descriptors. + static SOCKET FindListenSocket(ui16 port) { + TVector names; + TFsPath("/proc/self/fd").ListNames(names); + for (const auto& name : names) { + SOCKET socket = INVALID_SOCKET; + if (!TryFromString(name, socket)) { + continue; + } + int listening = 0; + socklen_t size = sizeof(listening); + if (getsockopt(socket, SOL_SOCKET, SO_ACCEPTCONN, &listening, &size) != 0 || !listening) { + continue; + } + sockaddr_in address{}; + size = sizeof(address); + if (getsockname(socket, reinterpret_cast(&address), &size) == 0 && + address.sin_family == AF_INET && InetToHost(address.sin_port) == port) + { + return socket; + } + } + return INVALID_SOCKET; + } + + // Existing yexception handlers must keep receiving errors, including their + // active exception context, while new consumers can recover the accept code. + static void CheckAcceptError(bool oneShot, size_t listenerThreads) { + class TCallback: public TEchoServer { + public: + TCallback() + : TEchoServer("ok") + { + } + + void OnAcceptException(int errorCode) override { + if (errorCode != EINVAL) { + ++MissingCode; + } + ++AcceptCalls; + THttpServer::ICallBack::OnAcceptException(errorCode); + } + + void OnException() override { + try { + throw; + } catch (const yexception& error) { + if (!TStringBuf(error.what()).Contains("accept:")) { + ++MissingContext; + } + } catch (...) { + ++MissingContext; + } + if (++Calls >= 2) { + Repeated.Signal(); + } + } + + std::atomic MissingCode = 0; + std::atomic MissingContext = 0; + std::atomic AcceptCalls = 0; + std::atomic Calls = 0; + TManualEvent Repeated; + } callback; + + TPortManager ports; + const ui16 port = ports.GetPort(); + THttpServer::TOptions options; + options.AddBindAddress("127.0.0.1", port); + options.OneShotPoll = oneShot; + options.nListenerThreads = listenerThreads; + THttpServer server(&callback, options); + UNIT_ASSERT(server.Start()); + const SOCKET listener = FindListenSocket(port); + UNIT_ASSERT_VALUES_UNEQUAL(listener, INVALID_SOCKET); + // Preserve the fd and poll registration but remove the LISTEN state. + UNIT_ASSERT_VALUES_EQUAL(shutdown(listener, SHUT_RDWR), 0); + const bool repeated = callback.Repeated.WaitT(TDuration::Seconds(5)); + server.Stop(); + UNIT_ASSERT_C(repeated, "The legacy callback retry policy must remain unchanged"); + UNIT_ASSERT_VALUES_EQUAL(callback.MissingCode.load(), 0); + UNIT_ASSERT_VALUES_EQUAL(callback.MissingContext.load(), 0); + UNIT_ASSERT_VALUES_EQUAL(callback.AcceptCalls.load(), callback.Calls.load()); + } + + Y_UNIT_TEST(AcceptErrorCodeLevelTriggered) { + CheckAcceptError(false, 1); + } + + Y_UNIT_TEST(AcceptErrorCodeOneShot) { + CheckAcceptError(true, 1); + } + + Y_UNIT_TEST(AcceptErrorCodeMultipleListeners) { + CheckAcceptError(true, 2); + } +#endif + } diff --git a/library/cpp/http/server/options.h b/library/cpp/http/server/options.h index d0b4f5d6234..9f5f0dc61a0 100644 --- a/library/cpp/http/server/options.h +++ b/library/cpp/http/server/options.h @@ -1,5 +1,7 @@ #pragma once +#include + #include #include #include @@ -46,6 +48,11 @@ class THttpServerOptions { return *this; } + inline THttpServerOptions& SetContentEncodingPredicate(TEncodeContentPredicate predicate) { + ContentEncodingPredicate = std::move(predicate); + return *this; + } + inline THttpServerOptions& EnableRejectExcessConnections(bool enable) noexcept { RejectExcessConnections = enable; @@ -164,6 +171,7 @@ class THttpServerOptions { bool KeepAliveEnabled = true; bool CompressionEnabled = false; + TEncodeContentPredicate ContentEncodingPredicate; bool RejectExcessConnections = false; bool ReusePort = false; // set SO_REUSEPORT socket option bool ReuseAddress = true; // set SO_REUSEADDR socket option diff --git a/library/cpp/logger/element_ut.cpp b/library/cpp/logger/element_ut.cpp index 86303973d8f..f7dc660641d 100644 --- a/library/cpp/logger/element_ut.cpp +++ b/library/cpp/logger/element_ut.cpp @@ -34,10 +34,10 @@ void TLogElementTest::TestMoveCtor() { THolder dst = MakeHolder(std::move(*src)); - src.Destroy(); + src.reset(); UNIT_ASSERT(output.Str() == ""); - dst.Destroy(); + dst.reset(); UNIT_ASSERT(output.Str() == message); } @@ -50,6 +50,6 @@ void TLogElementTest::TestWith() { TString message = "Hello, World!"; (*src).With("Foo", "Bar").With("Foo", "Baz") << message; - src.Destroy(); + src.reset(); UNIT_ASSERT(output.Str() == "Hello, World!; Foo=Bar; Foo=Baz; "); } diff --git a/library/cpp/logger/log.cpp b/library/cpp/logger/log.cpp index 286896a6ead..7f97278263c 100644 --- a/library/cpp/logger/log.cpp +++ b/library/cpp/logger/log.cpp @@ -9,6 +9,8 @@ #include #include +#include + THolder CreateLogBackend(const TString& fname, ELogPriority priority, bool threaded) { TLogBackendCreatorUninitialized creator; creator.InitCustom(fname, priority, threaded); @@ -90,7 +92,7 @@ class TLog::TImpl: public TAtomicRefCount { } inline void CloseLog() noexcept { - Backend_.Destroy(); + Backend_.reset(); Y_ASSERT(!IsOpen()); } @@ -137,6 +139,11 @@ TLog::TLog(THolder backend) { } +TLog::TLog(std::unique_ptr backend) + : TLog(THolder(backend.release())) +{ +} + TLog::TLog(const TLog&) = default; TLog::TLog(TLog&&) = default; TLog::~TLog() = default; @@ -215,6 +222,10 @@ void TLog::ResetBackend(THolder backend) noexcept { Impl_->ResetBackend(std::move(backend)); } +void TLog::ResetBackend(std::unique_ptr backend) noexcept { + ResetBackend(THolder(backend.release())); +} + bool TLog::IsNullLog() const noexcept { return Impl_->IsNullLog(); } diff --git a/library/cpp/logger/log.h b/library/cpp/logger/log.h index 6c90b0cb29c..a38ec365e4d 100644 --- a/library/cpp/logger/log.h +++ b/library/cpp/logger/log.h @@ -9,35 +9,37 @@ #include #include -#include #include +#include +#include using TLogFormatter = std::function; -// Logging facilities interface. -// -// ```cpp -// TLog base; -// ... -// auto log = base; -// log.SetFormatter([reqId](ELogPriority p, TStringBuf msg) { -// return TStringBuilder() << "reqid=" << reqId << "; " << msg; -// }); -// -// log.Write(TLOG_INFO, "begin"); -// HandleRequest(...); -// log.Write(TLOG_INFO, "end"); -// ``` -// -// Users are encouraged to copy `TLog` instance. +/// Logging facilities interface. +/// +/// @code +/// TLog base; +/// ... +/// auto log = base; +/// log.SetFormatter([reqId](ELogPriority p, TStringBuf msg) { +/// return TStringBuilder() << "reqid=" << reqId << "; " << msg; +/// }); +/// +/// log.Write(TLOG_INFO, "begin"); +/// HandleRequest(...); +/// log.Write(TLOG_INFO, "end"); +/// @endcode +/// +/// Users are encouraged to copy TLog instance. class TLog { public: - // Construct empty logger all writes will be spilled. + /// Construct empty logger all writes will be spilled. TLog(); - // Construct file logger. + /// Construct file logger. TLog(const TString& fname, ELogPriority priority = LOG_MAX_PRIORITY); - // Construct any type of logger + /// Construct any type of logger TLog(THolder backend); + TLog(std::unique_ptr backend); TLog(const TLog&); TLog(TLog&&); @@ -45,52 +47,56 @@ class TLog { TLog& operator=(const TLog&); TLog& operator=(TLog&&); - // Change underlying backend. - // NOTE: not thread safe. + /// Change underlying backend. + /// @note: not thread safe. + /// @{ void ResetBackend(THolder backend) noexcept; - // Reset underlying backend, `IsNullLog()` will return `true` after this call. - // NOTE: not thread safe. + void ResetBackend(std::unique_ptr backend) noexcept; + /// @} + + /// Reset underlying backend, IsNullLog() will return `true` after this call. + /// @note: not thread safe. THolder ReleaseBackend() noexcept; - // Check if underlying backend is defined and is not null. - // NOTE: not thread safe with respect to `ResetBackend` and `ReleaseBackend`. bool IsNullLog() const noexcept; + /// Check if underlying backend is defined and is not null. + /// @note: not thread safe with respect to ResetBackend() and ReleaseBackend(). bool IsNotNullLog() const noexcept { return !IsNullLog(); } - // Write message to the log. - // - // @param[in] priority Message priority to use. - // @param[in] message Message to write. - // @param[in] metaFlags Message meta flags. + /// Write message to the log. + /// + /// @param[in] priority Message priority to use. + /// @param[in] message Message to write. + /// @param[in] metaFlags Message meta flags. void Write(ELogPriority priority, TStringBuf message, TLogRecord::TMetaFlags metaFlags = {}) const; - // Write message to the log using `DefaultPriority()`. + /// Write message to the log using DefaultPriority(). void Write(const char* data, size_t len, TLogRecord::TMetaFlags metaFlags = {}) const; - // Write message to the log, but pass the message in a c-style. + /// Write message to the log, but pass the message in a c-style. void Write(ELogPriority priority, const char* data, size_t len, TLogRecord::TMetaFlags metaFlags = {}) const; - // Write message to the log in a c-like printf style. + /// Write message to the log in a c-like printf style. void Y_PRINTF_FORMAT(3, 4) AddLog(ELogPriority priority, const char* format, ...) const; - // Write message to the log in a c-like printf style with `DefaultPriority()` priority. + /// Write message to the log in a c-like printf style with DefaultPriority() priority. void Y_PRINTF_FORMAT(2, 3) AddLog(const char* format, ...) const; - // Call `ReopenLog()` of the underlying backend. + /// Call `ReopenLog()` of the underlying backend. void ReopenLog(); - // Call `ReopenLogNoFlush()` of the underlying backend. + /// Call `ReopenLogNoFlush()` of the underlying backend. void ReopenLogNoFlush(); - // Call `QueueSize()` of the underlying backend. + /// Call `QueueSize()` of the underlying backend. size_t BackEndQueueSize() const; - // Set log default priority. - // NOTE: not thread safe. + /// Set log default priority. + /// @note: not thread safe. void SetDefaultPriority(ELogPriority priority) noexcept; - // Get default priority + /// Get default priority ELogPriority DefaultPriority() const noexcept; - // Call `FiltrationLevel()` of the underlying backend. + /// Call `FiltrationLevel()` of the underlying backend. ELogPriority FiltrationLevel() const noexcept; - // Set current log formatter. + /// Set current log formatter. void SetFormatter(TLogFormatter formatter) noexcept; template @@ -101,8 +107,8 @@ class TLog { } public: - // These methods are deprecated and present here only for compatibility reasons (for 13 years - // already ...). Do not use them. + /// These methods are deprecated and present here only for compatibility reasons (for 13 years + /// already ...). Do not use them. bool OpenLog(const char* path, ELogPriority lp = LOG_MAX_PRIORITY); bool IsOpen() const noexcept; void AddLogVAList(const char* format, va_list lst); diff --git a/library/cpp/monlib/encode/CMakeLists.txt b/library/cpp/monlib/encode/CMakeLists.txt index 3ff87283038..817619ecb46 100644 --- a/library/cpp/monlib/encode/CMakeLists.txt +++ b/library/cpp/monlib/encode/CMakeLists.txt @@ -24,6 +24,7 @@ target_link_libraries(monlib-encode-json monlib-exception json json-writer + RapidJSON::RapidJSON ) target_sources(monlib-encode-json PRIVATE diff --git a/library/cpp/monlib/encode/buffered/buffered_encoder_base.cpp b/library/cpp/monlib/encode/buffered/buffered_encoder_base.cpp index c0449a10ff0..de3eaae9dd9 100644 --- a/library/cpp/monlib/encode/buffered/buffered_encoder_base.cpp +++ b/library/cpp/monlib/encode/buffered/buffered_encoder_base.cpp @@ -18,6 +18,11 @@ void TBufferedEncoderBase::OnCommonTime(TInstant time) { CommonTime_ = time; } +void TBufferedEncoderBase::OnCommonStartTimeSeconds(ui32 startTimeSeconds) { + State_.Expect(TEncoderState::EState::ROOT); + CommonStartTimeSeconds_ = startTimeSeconds; +} + void TBufferedEncoderBase::OnMetricBegin(EMetricType type) { State_.Switch(TEncoderState::EState::ROOT, TEncoderState::EState::METRIC); Metrics_.emplace_back(); @@ -150,6 +155,13 @@ void TBufferedEncoderBase::OnMemOnly(bool isMemOnly) { metric.IsMemOnly = isMemOnly; } +void TBufferedEncoderBase::OnStartTimeSeconds(ui32 startTimeSeconds) { + State_.Expect(TEncoderState::EState::METRIC); + TMetric& metric = Metrics_.back(); + metric.HasStartTime = true; + metric.StartTimeSeconds = startTimeSeconds; +} + TString TBufferedEncoderBase::FormatLabels(const TPooledLabels& labels) const { auto formattedLabels = TVector(Reserve(labels.size() + CommonLabels_.size())); auto addLabel = [&](const TPooledLabel& l) { diff --git a/library/cpp/monlib/encode/buffered/buffered_encoder_base.h b/library/cpp/monlib/encode/buffered/buffered_encoder_base.h index dab5671ad42..b550d33103a 100644 --- a/library/cpp/monlib/encode/buffered/buffered_encoder_base.h +++ b/library/cpp/monlib/encode/buffered/buffered_encoder_base.h @@ -19,6 +19,7 @@ class TBufferedEncoderBase : public IMetricEncoder { void OnStreamEnd() override; void OnCommonTime(TInstant time) override; + void OnCommonStartTimeSeconds(ui32 startTimeSeconds) override; void OnMetricBegin(EMetricType type) override; void OnMetricEnd() override; @@ -38,6 +39,7 @@ class TBufferedEncoderBase : public IMetricEncoder { void OnLogHistogram(TInstant, TLogHistogramSnapshotPtr) override; void OnMemOnly(bool isMemOnly) override; + void OnStartTimeSeconds(ui32 startTimeSeconds) override; protected: using TPooledStr = TStringPoolBuilder::TValue; @@ -82,7 +84,9 @@ class TBufferedEncoderBase : public IMetricEncoder { EMetricType MetricType = EMetricType::UNKNOWN; TPooledLabels Labels; TMetricTimeSeries TimeSeries; - bool IsMemOnly; + bool IsMemOnly = false; + bool HasStartTime = false; + ui32 StartTimeSeconds = 0; }; protected: @@ -94,6 +98,7 @@ class TBufferedEncoderBase : public IMetricEncoder { TStringPoolBuilder LabelNamesPool_; TStringPoolBuilder LabelValuesPool_; TInstant CommonTime_ = TInstant::Zero(); + ui32 CommonStartTimeSeconds_ = 0; TPooledLabels CommonLabels_; TVector Metrics_; TMetricMap MetricMap_; diff --git a/library/cpp/monlib/encode/format.cpp b/library/cpp/monlib/encode/format.cpp index dbcc2411f46..a43c0f8537d 100644 --- a/library/cpp/monlib/encode/format.cpp +++ b/library/cpp/monlib/encode/format.cpp @@ -7,6 +7,39 @@ #include namespace NMonitoring { + namespace { + constexpr TStringBuf SolomonMediaTypePrefix = "application/x-solomon-"; + + bool IsSolomonMediaType(TStringBuf value) { + value = StripString(value).Before(';'); + return value.size() > SolomonMediaTypePrefix.size() && + AsciiHasPrefixIgnoreCase(value, SolomonMediaTypePrefix); + } + + bool HeaderHasSolomonMediaType(const THttpHeaders& headers, TStringBuf headerName) { + for (const auto& header : headers) { + if (AsciiEqualsIgnoreCase(header.Name(), headerName) && IsSolomonMediaType(header.Value())) { + return true; + } + } + return false; + } + + bool AcceptsSolomonMediaType(const THttpHeaders& headers) { + for (const auto& header : headers) { + if (!AsciiEqualsIgnoreCase(header.Name(), "Accept")) { + continue; + } + for (const auto& item : StringSplitter(header.Value()).Split(',').SkipEmpty()) { + if (IsSolomonMediaType(item.Token())) { + return true; + } + } + } + return false; + } + } + static ECompression CompressionFromHeader(TStringBuf value) { if (value.empty()) { return ECompression::UNKNOWN; @@ -68,6 +101,11 @@ namespace NMonitoring { return FormatFromHttpMedia(value); } + bool DisableContentEncoding(const THttpHeaders& requestHeaders, const THttpHeaders& responseHeaders) { + return !AcceptsSolomonMediaType(requestHeaders) && + !HeaderHasSolomonMediaType(responseHeaders, "Content-Type"); + } + TStringBuf ContentTypeByFormat(EFormat format) { switch (format) { case EFormat::SPACK: diff --git a/library/cpp/monlib/encode/format.h b/library/cpp/monlib/encode/format.h index 08c59e192d9..ab464ee90fe 100644 --- a/library/cpp/monlib/encode/format.h +++ b/library/cpp/monlib/encode/format.h @@ -1,5 +1,7 @@ #pragma once +#include + #include namespace NMonitoring { @@ -133,6 +135,12 @@ namespace NMonitoring { */ EFormat FormatFromContentType(TStringBuf value); + /** + * Content-encoding policy that disables HTTP encoding for monitoring media types. + * Monitoring formats manage compression inside their payload and must pass through. + */ + bool DisableContentEncoding(const THttpHeaders& requestHeaders, const THttpHeaders& responseHeaders); + /** * Returns value for "Content-Type" header determined by the given * format type. diff --git a/library/cpp/monlib/encode/format_ut.cpp b/library/cpp/monlib/encode/format_ut.cpp index 92bcddc86bb..3051baf759e 100644 --- a/library/cpp/monlib/encode/format_ut.cpp +++ b/library/cpp/monlib/encode/format_ut.cpp @@ -67,6 +67,27 @@ Y_UNIT_TEST_SUITE(TFormatTest) { EFormat::PROMETHEUS); } + Y_UNIT_TEST(DisableContentEncodingForMonitoring) { + const THttpHeaders empty; + const THttpHeaders regularRequest({ + {"Accept", "application/json"}, + }); + const THttpHeaders solomonRequest({ + {"Accept", "application/json, application/x-solomon-spack"}, + }); + const THttpHeaders regularResponse({ + {"Content-Type", "application/json"}, + }); + const THttpHeaders solomonResponse({ + {"Content-Type", "Application/X-Solomon-Multi-Spack; version=1"}, + }); + + UNIT_ASSERT(DisableContentEncoding(empty, regularResponse)); + UNIT_ASSERT(DisableContentEncoding(regularRequest, regularResponse)); + UNIT_ASSERT(!DisableContentEncoding(solomonRequest, regularResponse)); + UNIT_ASSERT(!DisableContentEncoding(regularRequest, solomonResponse)); + } + Y_UNIT_TEST(FormatToStrFromStr) { const std::array formats = {{ EFormat::UNKNOWN, diff --git a/library/cpp/monlib/encode/json/json_decoder.cpp b/library/cpp/monlib/encode/json/json_decoder.cpp index 801ff0833d2..09cae18c434 100644 --- a/library/cpp/monlib/encode/json/json_decoder.cpp +++ b/library/cpp/monlib/encode/json/json_decoder.cpp @@ -8,6 +8,10 @@ #include +#include +#include +#include + #include #include @@ -435,6 +439,108 @@ class TCommonPartsProxy: public IHaltableMetricConsumer { bool IsMetric_{false}; }; +// TODO(SOLOMON-21639): Move the TStringBuf overload to library/cpp/json and remove this copy. +// Copied from library/cpp/json/json_reader.cpp because TJsonCallbacksWrapper is private to that +// translation unit. The optimization stays in monlib because downstream canons include source +// locations from the shared JSON library. +struct TJsonCallbacksWrapper { + NJson::TJsonCallbacks& Impl; + + explicit TJsonCallbacksWrapper(NJson::TJsonCallbacks& impl) + : Impl(impl) + { + } + + bool Null() { + return Impl.OnNull(); + } + + bool Bool(bool value) { + return Impl.OnBoolean(value); + } + + template + bool ProcessUint(T value) { + if (Y_LIKELY(value <= ui64(Max()))) { + return Impl.OnInteger(i64(value)); + } + return Impl.OnUInteger(value); + } + + bool Int(int value) { + return Impl.OnInteger(value); + } + + bool Uint(unsigned value) { + return ProcessUint(value); + } + + bool Int64(i64 value) { + return Impl.OnInteger(value); + } + + bool Uint64(ui64 value) { + return ProcessUint(value); + } + + bool Double(double value) { + return Impl.OnDouble(value); + } + + bool RawNumber(const char* value, rapidjson::SizeType size, bool copy) { + Y_ASSERT(false && "this method should never be called"); + Y_UNUSED(value); + Y_UNUSED(size); + Y_UNUSED(copy); + return true; + } + + bool String(const char* value, rapidjson::SizeType size, bool copy) { + Y_ASSERT(copy); + return Impl.OnString(TStringBuf(value, size)); + } + + bool StartObject() { + return Impl.OnOpenMap(); + } + + bool Key(const char* value, rapidjson::SizeType size, bool copy) { + Y_ASSERT(copy); + return Impl.OnMapKey(TStringBuf(value, size)); + } + + bool EndObject(rapidjson::SizeType memberCount) { + Y_UNUSED(memberCount); + return Impl.OnCloseMap(); + } + + bool StartArray() { + return Impl.OnOpenArray(); + } + + bool EndArray(rapidjson::SizeType elementCount) { + Y_UNUSED(elementCount); + return Impl.OnCloseArray(); + } +}; + +bool ReadJson(TStringBuf data, NJson::TJsonCallbacks* callbacks) { + TJsonCallbacksWrapper wrapper(*callbacks); + rapidjson::MemoryStream stream(data.data() ? data.data() : "", data.size()); + rapidjson::Reader reader; + auto result = reader.Parse(stream, wrapper); + + if (result.IsError()) { + auto reason = TStringBuilder() << "Offset: " << result.Offset() + << ", Code: " << static_cast(result.Code()) + << ", Error: " << rapidjson::GetParseError_En(result.Code()); + callbacks->OnError(result.Offset(), reason); + return false; + } + + return callbacks->OnEnd(); +} + /////////////////////////////////////////////////////////////////////// // TDecoderJson /////////////////////////////////////////////////////////////////////// @@ -1169,18 +1275,16 @@ if (Y_UNLIKELY(!(CONDITION))) { \ void DecodeJson(TStringBuf data, IMetricConsumer* c, TStringBuf metricNameLabel) { TCommonPartsCollector commonPartsCollector; { - TMemoryInput memIn(data); TDecoderJson decoder(data, &commonPartsCollector, metricNameLabel); // no need to check a return value. If there is an error, a TJsonDecodeError is thrown - NJson::ReadJson(&memIn, &decoder); + ReadJson(data, &decoder); } TCommonPartsProxy commonPartsProxy(std::move(commonPartsCollector.CommonParts()), c); { - TMemoryInput memIn(data); TDecoderJson decoder(data, &commonPartsProxy, metricNameLabel); // no need to check a return value. If there is an error, a TJsonDecodeError is thrown - NJson::ReadJson(&memIn, &decoder); + ReadJson(data, &decoder); } } diff --git a/library/cpp/monlib/encode/json/json_decoder_ut.cpp b/library/cpp/monlib/encode/json/json_decoder_ut.cpp index 4464e1d26a4..cf5aecbfd22 100644 --- a/library/cpp/monlib/encode/json/json_decoder_ut.cpp +++ b/library/cpp/monlib/encode/json/json_decoder_ut.cpp @@ -75,11 +75,10 @@ void ValidateMetrics(const TVector& metrics) { void CheckCommonPartsCollector(TString data, bool shouldBeStopped, bool checkLabels = true, bool checkTs = true, TStringBuf metricNameLabel = "name") { TCommonPartsCollector commonPartsCollector; - TMemoryInput memIn(data); TDecoderJson decoder(data, &commonPartsCollector, metricNameLabel); bool isOk{false}; - UNIT_ASSERT_NO_EXCEPTION(isOk = NJson::ReadJson(&memIn, &decoder)); + UNIT_ASSERT_NO_EXCEPTION(isOk = ReadJson(data, &decoder)); UNIT_ASSERT_VALUES_EQUAL(isOk, !shouldBeStopped); ValidateCommonParts(commonPartsCollector.CommonParts(), checkLabels, checkTs); diff --git a/library/cpp/monlib/encode/prometheus/prometheus_decoder_ut.cpp b/library/cpp/monlib/encode/prometheus/prometheus_decoder_ut.cpp index 51877c4699b..2d9b0e618ed 100644 --- a/library/cpp/monlib/encode/prometheus/prometheus_decoder_ut.cpp +++ b/library/cpp/monlib/encode/prometheus/prometheus_decoder_ut.cpp @@ -67,6 +67,17 @@ Y_UNIT_TEST_SUITE(TPrometheusDecoderTest) { } } + Y_UNIT_TEST(MetricNameWithDotsIsRejected) { + constexpr auto inputMetrics = + "# TYPE service_account.authorized_key.create_token_events_count_total counter\n" + "service_account.authorized_key.create_token_events_count_total 1\n"; + + UNIT_ASSERT_EXCEPTION_CONTAINS( + Decode(inputMetrics), + TPrometheusDecodeException, + "unknown metric type: .authorized_key.create_token_events_count_total"); + } + Y_UNIT_TEST(Minimal) { auto samples = Decode( "minimal_metric 1.234\n" diff --git a/library/cpp/monlib/encode/spack/spack_v1.h b/library/cpp/monlib/encode/spack/spack_v1.h index ce516626e31..d3d94265948 100644 --- a/library/cpp/monlib/encode/spack/spack_v1.h +++ b/library/cpp/monlib/encode/spack/spack_v1.h @@ -93,7 +93,8 @@ namespace NMonitoring { SV1_00 = 0x0100, SV1_01 = 0x0101, SV1_02 = 0x0102, - SV1_03 = 0x0103 + SV1_03 = 0x0103, + SV1_04 = 0x0104 }; IMetricEncoderPtr EncoderSpackV1( @@ -118,6 +119,13 @@ namespace NMonitoring { EMetricsMergingMode mergingMode = EMetricsMergingMode::DEFAULT ); + IMetricEncoderPtr EncoderSpackV14( + IOutputStream* out, + ETimePrecision timePrecision, + ECompression compression, + EMetricsMergingMode mergingMode = EMetricsMergingMode::DEFAULT + ); + void DecodeSpackV1(IInputStream* in, IMetricConsumer* c, TStringBuf metricNameLabel = "name"); } diff --git a/library/cpp/monlib/encode/spack/spack_v1_decoder.cpp b/library/cpp/monlib/encode/spack/spack_v1_decoder.cpp index 7d04aa0d555..a8f0d7bbf41 100644 --- a/library/cpp/monlib/encode/spack/spack_v1_decoder.cpp +++ b/library/cpp/monlib/encode/spack/spack_v1_decoder.cpp @@ -66,7 +66,7 @@ namespace NMonitoring { TVector namesBuf; TVector valuesBuf; - if (Header_.Version == SV1_03) { + if (Header_.Version == SV1_03 || Header_.Version == SV1_04) { auto namesResult = ReadLengthDelimitedStringPool(Header_.LabelNamesSize); namesBuf = std::move(namesResult.first); auto valuesResult = ReadLengthDelimitedStringPool(Header_.LabelValuesSize); @@ -107,6 +107,10 @@ namespace NMonitoring { // (3) read common time c->OnCommonTime(ReadTime()); + if (Header_.Version == SV1_04) { + c->OnCommonStartTimeSeconds(ReadFixed()); + } + // (4) read common labels if (ui32 commonLabelsCount = ReadVarint()) { c->OnLabelsBegin(); @@ -133,7 +137,11 @@ namespace NMonitoring { c->OnMetricBegin(metricType); // (5.2) flags byte - c->OnMemOnly(ReadFixed() & 0x01); + const ui8 flagsByte = ReadFixed(); + c->OnMemOnly(flagsByte & 0x01); + if (Header_.Version == SV1_04 && flagsByte & 0x02) { + c->OnStartTimeSeconds(ReadFixed()); + } auto metricNameValueIndex = std::numeric_limits::max(); if (Header_.Version == SV1_02) { diff --git a/library/cpp/monlib/encode/spack/spack_v1_encoder.cpp b/library/cpp/monlib/encode/spack/spack_v1_encoder.cpp index ac9934b0593..b0c41bd022e 100644 --- a/library/cpp/monlib/encode/spack/spack_v1_encoder.cpp +++ b/library/cpp/monlib/encode/spack/spack_v1_encoder.cpp @@ -6,6 +6,7 @@ #include #include +#include #include #ifndef _little_endian_ @@ -14,6 +15,17 @@ namespace NMonitoring { namespace { + ui32 DefaultStartTimeSeconds() noexcept { + static const ui32 startTimeSeconds = []() -> ui32 { + try { + return static_cast((TInstant::Now() - ProcessUptime()).Seconds()); + } catch (...) { + return 0; + } + }(); + return startTimeSeconds; + } + /////////////////////////////////////////////////////////////////////// // TEncoderSpackV1 /////////////////////////////////////////////////////////////////////// @@ -34,6 +46,9 @@ namespace NMonitoring { , MetricName_(Version_ == SV1_02 ? LabelNamesPool_.PutIfAbsent(metricNameLabel) : nullptr) { MetricsMergingMode_ = mergingMode; + if (Version_ == SV1_04) { + CommonStartTimeSeconds_ = DefaultStartTimeSeconds(); + } LabelNamesPool_.SetSorted(true); LabelValuesPool_.SetSorted(true); @@ -96,7 +111,7 @@ namespace NMonitoring { header.Version = Version_; header.TimePrecision = EncodeTimePrecision(TimePrecision_); header.Compression = EncodeCompression(Compression_); - if (Version_ == SV1_03) { + if (Version_ == SV1_03 || Version_ == SV1_04) { header.LabelNamesSize = static_cast(LabelNamesPool_.Count()); header.LabelValuesSize = static_cast(LabelValuesPool_.Count()); } else { @@ -116,7 +131,7 @@ namespace NMonitoring { } // (2) write string pools - if (Version_ == SV1_03) { + if (Version_ == SV1_03 || Version_ == SV1_04) { auto strPoolWrite = [this](TStringBuf str, ui32, ui32) { WriteVarUInt32(Out_, static_cast(str.size())); Out_->Write(str); @@ -135,6 +150,10 @@ namespace NMonitoring { // (3) write common time WriteTime(CommonTime_); + if (Version_ == SV1_04) { + WriteFixed(CommonStartTimeSeconds_); + } + // (4) write common labels' indexes WriteLabels(CommonLabels_, nullptr); @@ -146,9 +165,17 @@ namespace NMonitoring { Out_->Write(&typesByte, sizeof(typesByte)); // (5.2) flags byte - ui8 flagsByte = metric.IsMemOnly & 0x01; + const bool writeStartTime = Version_ == SV1_04 + && metric.HasStartTime + && metric.StartTimeSeconds != CommonStartTimeSeconds_ + && (metric.MetricType == EMetricType::RATE || metric.MetricType == EMetricType::HIST_RATE); + ui8 flagsByte = (metric.IsMemOnly & 0x01) | (static_cast(writeStartTime) << 1); Out_->Write(&flagsByte, sizeof(flagsByte)); + if (writeStartTime) { + WriteFixed(metric.StartTimeSeconds); + } + // v1.2 format addition — metric name if (Version_ == SV1_02) { const auto it = FindIf(metric.Labels, [&](const auto& l) { @@ -341,4 +368,13 @@ namespace NMonitoring { ) { return MakeHolder(out, timePrecision, compression, mergingMode, SV1_03, ""); } + + IMetricEncoderPtr EncoderSpackV14( + IOutputStream* out, + ETimePrecision timePrecision, + ECompression compression, + EMetricsMergingMode mergingMode + ) { + return MakeHolder(out, timePrecision, compression, mergingMode, SV1_04, ""); + } } diff --git a/library/cpp/monlib/encode/spack/spack_v1_ut.cpp b/library/cpp/monlib/encode/spack/spack_v1_ut.cpp index 2a75f441f2d..8d38d1f61e7 100644 --- a/library/cpp/monlib/encode/spack/spack_v1_ut.cpp +++ b/library/cpp/monlib/encode/spack/spack_v1_ut.cpp @@ -1,5 +1,6 @@ #include "spack_v1.h" +#include #include #include #include @@ -49,6 +50,16 @@ void AssertPointEqual(const NProto::TPoint& p, TInstant time, i64 value) { UNIT_ASSERT_VALUES_EQUAL(p.GetInt64(), value); } +class TStartTimeCollectingConsumer final: public TCollectingConsumer { +public: + void OnStartTimeSeconds(ui32 startTimeSeconds) override { + StartTimeSeconds.push_back(startTimeSeconds); + TCollectingConsumer::OnStartTimeSeconds(startTimeSeconds); + } + + TVector StartTimeSeconds; +}; + Y_UNIT_TEST_SUITE(TSpackTest) { ui8 expectedHeader_v1_0[] = { 0x53, 0x50, // magic "SP" (fixed ui16) @@ -1223,6 +1234,135 @@ Y_UNIT_TEST_SUITE(TSpackTest) { UNIT_ASSERT_VALUES_EQUAL(header->PointsCount, 0u); } + Y_UNIT_TEST(V14StartTimeSeconds) { + constexpr ui32 commonStartTimeSeconds = 1'700'000'000; + constexpr ui32 metricStartTimeSeconds = 1'700'000'100; + constexpr ui32 histogramStartTimeSeconds = 1'700'000'200; + const TInstant commonTime = TInstant::Seconds(commonStartTimeSeconds + 15); + + TBuffer buffer; + { + TBufferOutput out(buffer); + auto e = EncoderSpackV14(&out, ETimePrecision::SECONDS, ECompression::IDENTITY); + + auto writeLabels = [&](TStringBuf sensor) { + e->OnLabelsBegin(); + e->OnLabel("sensor", sensor); + e->OnLabelsEnd(); + }; + + e->OnStreamBegin(); + e->OnCommonTime(commonTime); + e->OnCommonStartTimeSeconds(commonStartTimeSeconds); + + e->OnMetricBegin(EMetricType::RATE); + writeLabels("shared"); + TRate{10, commonStartTimeSeconds}.Accept(commonTime, e.Get()); + e->OnMetricEnd(); + + e->OnMetricBegin(EMetricType::GAUGE); + writeLabels("noStart"); + e->OnStartTimeSeconds(metricStartTimeSeconds + 1); + e->OnDouble(commonTime, 1.5); + e->OnMetricEnd(); + + e->OnMetricBegin(EMetricType::RATE); + writeLabels("zero"); + TRate{1, 0}.Accept(commonTime, e.Get()); + e->OnMetricEnd(); + + e->OnMetricBegin(EMetricType::RATE); + writeLabels("own"); + TLazyRate{[] { return ui64{20}; }, metricStartTimeSeconds}.Accept( + TInstant::Seconds(metricStartTimeSeconds + 5), + e.Get()); + e->OnMetricEnd(); + + e->OnMetricBegin(EMetricType::HIST_RATE); + writeLabels("histogram"); + e->OnStartTimeSeconds(histogramStartTimeSeconds); + e->OnHistogram( + TInstant::Seconds(histogramStartTimeSeconds + 5), + ExplicitHistogramSnapshot({10, 20, Max()}, {1, 2, 3})); + e->OnMetricEnd(); + + e->OnStreamEnd(); + e->Close(); + } + + const auto* header = reinterpret_cast(buffer.Data()); + UNIT_ASSERT_VALUES_EQUAL(header->Version, static_cast(SV1_04)); + UNIT_ASSERT_VALUES_EQUAL(header->MetricCount, 5u); + UNIT_ASSERT_VALUES_EQUAL(header->PointsCount, 5u); + + TStartTimeCollectingConsumer consumer; + TBufferInput in(buffer); + DecodeSpackV1(&in, &consumer); + + UNIT_ASSERT_VALUES_EQUAL(consumer.CommonStartTimeSeconds, commonStartTimeSeconds); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics.size(), 5u); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[0].StartTimeSeconds, 0u); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[1].StartTimeSeconds, 0u); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[2].StartTimeSeconds, 0u); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[3].StartTimeSeconds, metricStartTimeSeconds); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[4].StartTimeSeconds, histogramStartTimeSeconds); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds.size(), 3u); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds[0], 0u); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds[1], metricStartTimeSeconds); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds[2], histogramStartTimeSeconds); + } + + Y_UNIT_TEST(V14DecoderForwardsStartTimeFlagForNonRate) { + constexpr ui32 startTimeSeconds = 1'700'000'100; + TSpackHeader header; + header.Version = SV1_04; + header.TimePrecision = EncodeTimePrecision(ETimePrecision::SECONDS); + header.Compression = EncodeCompression(ECompression::IDENTITY); + header.LabelNamesSize = 1; + header.LabelValuesSize = 1; + header.MetricCount = 1; + header.PointsCount = 0; + + TBuffer buffer; + TBufferOutput out(buffer); + out.Write(&header, sizeof(header)); + + constexpr TStringBuf labelName = "sensor"; + constexpr TStringBuf labelValue = "gauge"; + const ui8 labelNameSize = labelName.size(); + const ui8 labelValueSize = labelValue.size(); + out.Write(&labelNameSize, sizeof(labelNameSize)); + out.Write(labelName.data(), labelName.size()); + out.Write(&labelValueSize, sizeof(labelValueSize)); + out.Write(labelValue.data(), labelValue.size()); + + constexpr ui32 commonTime = 0; + constexpr ui32 commonStartTimeSeconds = 0; + constexpr ui8 noLabels = 0; + constexpr ui8 oneLabel = 1; + constexpr ui8 typesByte = EncodeMetricType(EMetricType::GAUGE) << 2; + constexpr ui8 flagsByte = 0x02; + out.Write(&commonTime, sizeof(commonTime)); + out.Write(&commonStartTimeSeconds, sizeof(commonStartTimeSeconds)); + out.Write(&noLabels, sizeof(noLabels)); + out.Write(&typesByte, sizeof(typesByte)); + out.Write(&flagsByte, sizeof(flagsByte)); + out.Write(&startTimeSeconds, sizeof(startTimeSeconds)); + out.Write(&oneLabel, sizeof(oneLabel)); + out.Write(&noLabels, sizeof(noLabels)); + out.Write(&noLabels, sizeof(noLabels)); + + TStartTimeCollectingConsumer consumer; + TBufferInput in(buffer); + DecodeSpackV1(&in, &consumer); + + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds.size(), 1u); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds[0], startTimeSeconds); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics.size(), 1u); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[0].Kind, EMetricType::GAUGE); + UNIT_ASSERT_VALUES_EQUAL(consumer.Metrics[0].StartTimeSeconds, startTimeSeconds); + } + Y_UNIT_TEST(V12MissingNameForOneMetric) { TBuffer b; TBufferOutput out(b); diff --git a/library/cpp/monlib/metrics/fake.h b/library/cpp/monlib/metrics/fake.h index a058e1d99a1..9abb5f9e446 100644 --- a/library/cpp/monlib/metrics/fake.h +++ b/library/cpp/monlib/metrics/fake.h @@ -96,6 +96,10 @@ namespace NMonitoring { return 0; } + ui32 StartTimeSeconds() const noexcept override { + return 0; + } + void Reset() noexcept override { } }; @@ -105,6 +109,10 @@ namespace NMonitoring { return 0; } + ui32 StartTimeSeconds() const noexcept override { + return 0; + } + void Reset() noexcept override {} }; diff --git a/library/cpp/monlib/metrics/metric.h b/library/cpp/monlib/metrics/metric.h index bb5fda322ea..44855f5d64b 100644 --- a/library/cpp/monlib/metrics/metric.h +++ b/library/cpp/monlib/metrics/metric.h @@ -108,6 +108,9 @@ namespace NMonitoring { virtual ui64 Add(ui64 n) noexcept = 0; virtual ui64 Get() const noexcept = 0; + virtual ui32 StartTimeSeconds() const noexcept { + return 0; + } }; class ILazyRate: public IMetric { @@ -117,6 +120,9 @@ namespace NMonitoring { } virtual ui64 Get() const noexcept = 0; + virtual ui32 StartTimeSeconds() const noexcept { + return 0; + } }; class IHistogram: public IMetric { @@ -133,6 +139,9 @@ namespace NMonitoring { virtual void Record(double value) noexcept = 0; virtual void Record(double value, ui32 count) noexcept = 0; virtual IHistogramSnapshotPtr TakeSnapshot() const = 0; + virtual ui32 StartTimeSeconds() const noexcept { + return 0; + } protected: const bool IsRate_; @@ -309,7 +318,9 @@ namespace NMonitoring { /////////////////////////////////////////////////////////////////////////////// class TRate final: public IRate { public: - explicit TRate(ui64 value = 0) { + explicit TRate(ui64 value = 0, ui32 startTimeSeconds = 0) + : StartTimeSeconds_{startTimeSeconds} + { Value_.store(value, std::memory_order_relaxed); } @@ -321,16 +332,22 @@ namespace NMonitoring { return Value_.load(std::memory_order_relaxed); } + ui32 StartTimeSeconds() const noexcept override { + return StartTimeSeconds_; + } + void Reset() noexcept override { Value_.store(0, std::memory_order_relaxed); } void Accept(TInstant time, IMetricConsumer* consumer) const override { + consumer->OnStartTimeSeconds(StartTimeSeconds_); consumer->OnUint64(time, Get()); } private: std::atomic_uint64_t Value_; + ui32 StartTimeSeconds_; }; /////////////////////////////////////////////////////////////////////////////// @@ -338,8 +355,9 @@ namespace NMonitoring { /////////////////////////////////////////////////////////////////////////////// class TLazyRate final: public ILazyRate { public: - explicit TLazyRate(std::function supplier) + explicit TLazyRate(std::function supplier, ui32 startTimeSeconds = 0) : Supplier_(std::move(supplier)) + , StartTimeSeconds_{startTimeSeconds} { } @@ -347,7 +365,12 @@ namespace NMonitoring { return Supplier_(); } + ui32 StartTimeSeconds() const noexcept override { + return StartTimeSeconds_; + } + void Accept(TInstant time, IMetricConsumer* consumer) const override { + consumer->OnStartTimeSeconds(StartTimeSeconds_); consumer->OnUint64(time, Get()); } @@ -355,6 +378,7 @@ namespace NMonitoring { private: std::function Supplier_; + ui32 StartTimeSeconds_; }; /////////////////////////////////////////////////////////////////////////////// @@ -365,12 +389,14 @@ namespace NMonitoring { THistogram(IHistogramCollectorPtr collector, bool isRate) : IHistogram(isRate) , Collector_(std::move(collector)) + , StartTimeSeconds_(isRate ? static_cast(TInstant::Now().Seconds()) : 0) { } THistogram(std::function makeHistogramCollector, bool isRate) : IHistogram(isRate) , Collector_(makeHistogramCollector()) + , StartTimeSeconds_(isRate ? static_cast(TInstant::Now().Seconds()) : 0) { } @@ -383,9 +409,16 @@ namespace NMonitoring { } void Accept(TInstant time, IMetricConsumer* consumer) const override { + if (IsRate_) { + consumer->OnStartTimeSeconds(StartTimeSeconds_); + } consumer->OnHistogram(time, TakeSnapshot()); } + ui32 StartTimeSeconds() const noexcept override { + return StartTimeSeconds_; + } + IHistogramSnapshotPtr TakeSnapshot() const override { return Collector_->Snapshot(); } @@ -396,5 +429,6 @@ namespace NMonitoring { private: IHistogramCollectorPtr Collector_; + ui32 StartTimeSeconds_; }; } diff --git a/library/cpp/monlib/metrics/metric_consumer.h b/library/cpp/monlib/metrics/metric_consumer.h index 1152b51c4af..428402a26da 100644 --- a/library/cpp/monlib/metrics/metric_consumer.h +++ b/library/cpp/monlib/metrics/metric_consumer.h @@ -16,6 +16,9 @@ namespace NMonitoring { virtual void OnStreamEnd() = 0; virtual void OnCommonTime(TInstant time) = 0; + virtual void OnCommonStartTimeSeconds(ui32 startTimeSeconds) { + Y_UNUSED(startTimeSeconds); + } virtual void OnMetricBegin(EMetricType type) = 0; virtual void OnMetricEnd() = 0; @@ -37,6 +40,10 @@ namespace NMonitoring { virtual void OnMemOnly(bool isMemOnly) { Y_UNUSED(isMemOnly); } + + virtual void OnStartTimeSeconds(ui32 startTimeSeconds) { + Y_UNUSED(startTimeSeconds); + } }; using IMetricConsumerPtr = THolder; diff --git a/library/cpp/monlib/metrics/metric_registry.cpp b/library/cpp/monlib/metrics/metric_registry.cpp index fa8d8b8ec5d..e06c6836842 100644 --- a/library/cpp/monlib/metrics/metric_registry.cpp +++ b/library/cpp/monlib/metrics/metric_registry.cpp @@ -143,28 +143,44 @@ namespace NMonitoring { return RateWithOpts(std::move(labels)); } TRate* TMetricRegistry::RateWithOpts(TLabels labels, TMetricOpts opts) { - return Metric(std::move(labels), std::move(opts)); + return Metric( + std::move(labels), + std::move(opts), + ui64{0}, + static_cast(TInstant::Now().Seconds())); } TRate* TMetricRegistry::Rate(ILabelsPtr labels) { return RateWithOpts(std::move(labels)); } TRate* TMetricRegistry::RateWithOpts(ILabelsPtr labels, TMetricOpts opts) { - return Metric(std::move(labels), std::move(opts)); + return Metric( + std::move(labels), + std::move(opts), + ui64{0}, + static_cast(TInstant::Now().Seconds())); } TLazyRate* TMetricRegistry::LazyRate(TLabels labels, std::function supplier) { return LazyRateWithOpts(std::move(labels), std::move(supplier)); } TLazyRate* TMetricRegistry::LazyRateWithOpts(TLabels labels, std::function supplier, TMetricOpts opts) { - return Metric(std::move(labels), std::move(opts), std::move(supplier)); + return Metric( + std::move(labels), + std::move(opts), + std::move(supplier), + static_cast(TInstant::Now().Seconds())); } TLazyRate* TMetricRegistry::LazyRate(ILabelsPtr labels, std::function supplier) { return LazyRateWithOpts(std::move(labels), std::move(supplier)); } TLazyRate* TMetricRegistry::LazyRateWithOpts(ILabelsPtr labels, std::function supplier, TMetricOpts opts) { - return Metric(std::move(labels), std::move(opts), std::move(supplier)); + return Metric( + std::move(labels), + std::move(opts), + std::move(supplier), + static_cast(TInstant::Now().Seconds())); } THistogram* TMetricRegistry::HistogramCounter(TLabels labels, IHistogramCollectorPtr collector) { @@ -242,6 +258,7 @@ namespace NMonitoring { void TMetricRegistry::Accept(TInstant time, IMetricConsumer* consumer) const { consumer->OnStreamBegin(); + consumer->OnCommonStartTimeSeconds(StartTimeSeconds_); if (!CommonLabels_.Empty()) { consumer->OnLabelsBegin(); @@ -307,6 +324,7 @@ namespace NMonitoring { void TMetricRegistry::Took(TInstant time, IMetricConsumer* consumer) const { consumer->OnStreamBegin(); + consumer->OnCommonStartTimeSeconds(StartTimeSeconds_); TVector> tmpMetrics; diff --git a/library/cpp/monlib/metrics/metric_registry.h b/library/cpp/monlib/metrics/metric_registry.h index 944bee1c6ac..68cee059ba8 100644 --- a/library/cpp/monlib/metrics/metric_registry.h +++ b/library/cpp/monlib/metrics/metric_registry.h @@ -315,6 +315,7 @@ namespace NMonitoring { THashMap Metrics_; TLabels CommonLabels_; + ui32 StartTimeSeconds_ = static_cast(TInstant::Now().Seconds()); }; void WriteLabels(IMetricConsumer* consumer, const ILabels& labels); diff --git a/library/cpp/monlib/metrics/metric_registry_ut.cpp b/library/cpp/monlib/metrics/metric_registry_ut.cpp index 89e6d56146f..5c3173745ef 100644 --- a/library/cpp/monlib/metrics/metric_registry_ut.cpp +++ b/library/cpp/monlib/metrics/metric_registry_ut.cpp @@ -1,5 +1,6 @@ #include "metric_registry.h" +#include #include #include #include @@ -38,6 +39,16 @@ void Out(IOutputStream& os, NMoni } namespace { + class TStartTimeCollectingConsumer final: public TCollectingConsumer { + public: + void OnStartTimeSeconds(ui32 startTimeSeconds) override { + StartTimeSeconds.push_back(startTimeSeconds); + TCollectingConsumer::OnStartTimeSeconds(startTimeSeconds); + } + + TVector StartTimeSeconds; + }; + template auto EnsureIdempotent(F&& f) { auto firstResult = f(); @@ -158,11 +169,18 @@ Y_UNIT_TEST_SUITE(TMetricRegistryTest) { TMetricRegistry registry(TLabels{{"common", "label"}}); ui64 val = 0; + const auto beforeCreation = static_cast(TInstant::Now().Seconds()); TLazyRate* r = EnsureIdempotent([&] { return registry.LazyRate({{"my", "rate"}}, [&val](){return val;}); }); + const auto afterCreation = static_cast(TInstant::Now().Seconds()); UNIT_ASSERT_VALUES_EQUAL(r->Get(), 0); + UNIT_ASSERT_GE(r->StartTimeSeconds(), beforeCreation); + UNIT_ASSERT_LE(r->StartTimeSeconds(), afterCreation); val = 42; UNIT_ASSERT_VALUES_EQUAL(r->Get(), 42); + + TLazyRate rateWithExplicitStartTime{[&val](){return val;}, 42}; + UNIT_ASSERT_VALUES_EQUAL(rateWithExplicitStartTime.StartTimeSeconds(), 42); } Y_UNIT_TEST(DoubleCounter) { @@ -468,4 +486,35 @@ Y_UNIT_TEST_SUITE(TMetricRegistryTest) { ExponentialHistogram(5, 2)), yexception); } + + Y_UNIT_TEST(RateStartTimeSeconds) { + TMetricRegistry registry; + const auto beforeCreation = static_cast(TInstant::Now().Seconds()); + auto* rate = registry.Rate({{"some", "rate"}}); + auto* histogramRate = registry.HistogramRate( + {{"some", "histogram_rate"}}, + ExponentialHistogram(5, 2)); + registry.HistogramCounter( + {{"some", "histogram_counter"}}, + ExponentialHistogram(5, 2)); + const auto afterCreation = static_cast(TInstant::Now().Seconds()); + + UNIT_ASSERT_GE(rate->StartTimeSeconds(), beforeCreation); + UNIT_ASSERT_LE(rate->StartTimeSeconds(), afterCreation); + UNIT_ASSERT_VALUES_EQUAL(registry.Rate({{"some", "rate"}})->StartTimeSeconds(), rate->StartTimeSeconds()); + UNIT_ASSERT_GE(histogramRate->StartTimeSeconds(), beforeCreation); + UNIT_ASSERT_LE(histogramRate->StartTimeSeconds(), afterCreation); + UNIT_ASSERT_VALUES_EQUAL( + registry.HistogramRate({{"some", "histogram_rate"}}, ExponentialHistogram(5, 2))->StartTimeSeconds(), + histogramRate->StartTimeSeconds()); + + TStartTimeCollectingConsumer consumer; + registry.Accept(TInstant::Now(), &consumer); + UNIT_ASSERT_VALUES_EQUAL(consumer.StartTimeSeconds.size(), 2u); + UNIT_ASSERT(Find(consumer.StartTimeSeconds, rate->StartTimeSeconds()) != consumer.StartTimeSeconds.end()); + UNIT_ASSERT(Find(consumer.StartTimeSeconds, histogramRate->StartTimeSeconds()) != consumer.StartTimeSeconds.end()); + + TRate rateWithExplicitStartTime{0, 42}; + UNIT_ASSERT_VALUES_EQUAL(rateWithExplicitStartTime.StartTimeSeconds(), 42); + } } diff --git a/library/cpp/monlib/service/monservice.cpp b/library/cpp/monlib/service/monservice.cpp index 124b67c72bb..fb22e667af4 100644 --- a/library/cpp/monlib/service/monservice.cpp +++ b/library/cpp/monlib/service/monservice.cpp @@ -12,12 +12,7 @@ using namespace NMonitoring; -TMonService2::TMonService2(ui16 port, const TString& host, ui32 threads, const TString& title, THolder auth) - : TMonService2(HttpServerOptions(port, host, threads), title, std::move(auth)) -{ -} - -TMonService2::TMonService2(const THttpServerOptions& options, const TString& title, THolder auth) +TMonService2::TMonService2(std::unique_ptr auth, const THttpServerOptions& options, const TString& title) : NMonitoring::TMtHttpServer(options, std::bind(&TMonService2::ServeRequest, this, std::placeholders::_1, std::placeholders::_2)) , Title(title) , IndexMonPage(new TIndexMonPage("", Title)) @@ -28,7 +23,7 @@ TMonService2::TMonService2(const THttpServerOptions& options, const TString& tit ctime_r(&t, StartTime); } -TMonService2::TMonService2(const THttpServerOptions& options, TSimpleSharedPtr pool, const TString& title, THolder auth) +TMonService2::TMonService2(std::unique_ptr auth, const THttpServerOptions& options, TSimpleSharedPtr pool, const TString& title) : NMonitoring::TMtHttpServer(options, std::bind(&TMonService2::ServeRequest, this, std::placeholders::_1, std::placeholders::_2), std::move(pool)) , Title(title) , IndexMonPage(new TIndexMonPage("", Title)) @@ -39,13 +34,28 @@ TMonService2::TMonService2(const THttpServerOptions& options, TSimpleSharedPtr auth) + : TMonService2(port, TString(), 0, title, std::move(auth)) +{ +} + TMonService2::TMonService2(ui16 port, ui32 threads, const TString& title, THolder auth) : TMonService2(port, TString(), threads, title, std::move(auth)) { } -TMonService2::TMonService2(ui16 port, const TString& title, THolder auth) - : TMonService2(port, TString(), 0, title, std::move(auth)) +TMonService2::TMonService2(ui16 port, const TString& host, ui32 threads, const TString& title, THolder auth) + : TMonService2(std::unique_ptr(auth.Release()), HttpServerOptions(port, host, threads), title) +{ +} + +TMonService2::TMonService2(const THttpServerOptions& options, const TString& title, THolder auth) + : TMonService2(std::unique_ptr(auth.Release()), options, title) +{ +} + +TMonService2::TMonService2(const THttpServerOptions& options, TSimpleSharedPtr pool, const TString& title, THolder auth) + : TMonService2(std::unique_ptr(auth.Release()), options, std::move(pool), title) { } diff --git a/library/cpp/monlib/service/monservice.h b/library/cpp/monlib/service/monservice.h index 299b9a48a4f..8c1f590d6d6 100644 --- a/library/cpp/monlib/service/monservice.h +++ b/library/cpp/monlib/service/monservice.h @@ -10,6 +10,7 @@ #include #include +#include namespace NMonitoring { class TMonService2: public TMtHttpServer { @@ -17,7 +18,7 @@ namespace NMonitoring { const TString Title; char StartTime[26]; TIntrusivePtr IndexMonPage; - THolder AuthProvider_; + std::unique_ptr AuthProvider_; public: static THttpServerOptions HttpServerOptions(ui16 port, const TString& host, ui32 threads) { @@ -44,6 +45,10 @@ namespace NMonitoring { explicit TMonService2(const THttpServerOptions& options, const TString& title = GetProgramName(), THolder auth = nullptr); explicit TMonService2(const THttpServerOptions& options, TSimpleSharedPtr pool, const TString& title = GetProgramName(), THolder auth = nullptr); + // Auth comes first to keep legacy nullptr and {} calls unambiguous. + explicit TMonService2(std::unique_ptr auth, const THttpServerOptions& options, const TString& title = GetProgramName()); + explicit TMonService2(std::unique_ptr auth, const THttpServerOptions& options, TSimpleSharedPtr pool, const TString& title = GetProgramName()); + ~TMonService2() override { Stop(); } diff --git a/library/cpp/monlib/service/service.cpp b/library/cpp/monlib/service/service.cpp index 93b2d10b0c9..f2c93f3e2f9 100644 --- a/library/cpp/monlib/service/service.cpp +++ b/library/cpp/monlib/service/service.cpp @@ -4,6 +4,7 @@ #include #include #include +#include #include #include @@ -13,6 +14,16 @@ #include namespace NMonitoring { + THttpServerOptions WithMonitoringContentEncodingPolicy(THttpServerOptions options) { + auto previousPredicate = std::move(options.ContentEncodingPredicate); + options.SetContentEncodingPredicate( + [previousPredicate = std::move(previousPredicate)](const THttpHeaders& requestHeaders, const THttpHeaders& responseHeaders) { + return (!previousPredicate || previousPredicate(requestHeaders, responseHeaders)) && + DisableContentEncoding(requestHeaders, responseHeaders); + }); + return options; + } + class THttpClient: public IHttpRequest { public: void ServeRequest(THttpInput& in, IOutputStream& out, const NAddr::IRemoteAddr* remoteAddr, const THandler& Handler) { @@ -117,6 +128,7 @@ namespace NMonitoring { TContIO io(Socket, c); THttpInput in(&io); THttpOutput out(&io, &in); + out.SetContentEncodingPredicate(DisableContentEncoding); // buffer reply so there will be ne context switching TStringStream s; ServeRequest(in, s, RemoteAddr, Parent.Handler); @@ -204,13 +216,13 @@ namespace NMonitoring { }; TMtHttpServer::TMtHttpServer(const TOptions& options, THandler handler, IThreadFactory* pool) - : THttpServer(this, options, pool) + : THttpServer(this, WithMonitoringContentEncodingPolicy(options), pool) , Handler(std::move(handler)) { } TMtHttpServer::TMtHttpServer(const TOptions& options, THandler handler, TSimpleSharedPtr pool) - : THttpServer(this, /* mainWorkers = */pool, /* failWorkers = */pool, options) + : THttpServer(this, /* mainWorkers = */pool, /* failWorkers = */pool, WithMonitoringContentEncodingPolicy(options)) , Handler(std::move(handler)) { } diff --git a/library/cpp/monlib/service/service.h b/library/cpp/monlib/service/service.h index dae66cad8a4..ec63cdc92ee 100644 --- a/library/cpp/monlib/service/service.h +++ b/library/cpp/monlib/service/service.h @@ -13,6 +13,10 @@ struct TMonitor; namespace NMonitoring { + /// Adds the monitoring media-type content-encoding policy while preserving + /// any policy already configured by the caller. + THttpServerOptions WithMonitoringContentEncodingPolicy(THttpServerOptions options); + struct IHttpRequest { virtual ~IHttpRequest() { } diff --git a/library/cpp/openssl/crypto/CMakeLists.txt b/library/cpp/openssl/crypto/CMakeLists.txt new file mode 100644 index 00000000000..8045486a7be --- /dev/null +++ b/library/cpp/openssl/crypto/CMakeLists.txt @@ -0,0 +1,6 @@ +_ydb_sdk_add_library(openssl-crypto) + +target_sources(openssl-crypto PRIVATE sha.cpp) +target_link_libraries(openssl-crypto PUBLIC yutil OpenSSL::Crypto) + +_ydb_sdk_install_targets(TARGETS openssl-crypto) diff --git a/library/cpp/openssl/crypto/sha.cpp b/library/cpp/openssl/crypto/sha.cpp new file mode 100644 index 00000000000..ffd999a238a --- /dev/null +++ b/library/cpp/openssl/crypto/sha.cpp @@ -0,0 +1,89 @@ +#include "sha.h" + +#include + +#include + +namespace NOpenSsl { + namespace NSha1 { + static_assert(DIGEST_LENGTH == SHA_DIGEST_LENGTH); + + TDigest Calc(const void* data, size_t dataSize) { + TDigest digest; + Y_ENSURE(SHA1(static_cast(data), dataSize, digest.data()) != nullptr); + return digest; + } + + TCalcer::TCalcer() + : Context{new SHAstate_st} { + Y_ENSURE(SHA1_Init(Context.Get()) == 1); + } + + TCalcer::~TCalcer() { + } + + void TCalcer::Update(const void* data, size_t dataSize) { + Y_ENSURE(SHA1_Update(Context.Get(), data, dataSize) == 1); + } + + TDigest TCalcer::Final() { + TDigest digest; + Y_ENSURE(SHA1_Final(digest.data(), Context.Get()) == 1); + return digest; + } + } + namespace NSha256 { + static_assert(DIGEST_LENGTH == SHA256_DIGEST_LENGTH); + + TDigest Calc(const void* data, size_t dataSize) { + TDigest digest; + Y_ENSURE(SHA256(static_cast(data), dataSize, digest.data()) != nullptr); + return digest; + } + + TCalcer::TCalcer() + : Context{new SHA256state_st} { + Y_ENSURE(SHA256_Init(Context.Get()) == 1); + } + + TCalcer::~TCalcer() { + } + + void TCalcer::Update(const void* data, size_t dataSize) { + Y_ENSURE(SHA256_Update(Context.Get(), data, dataSize) == 1); + } + + TDigest TCalcer::Final() { + TDigest digest; + Y_ENSURE(SHA256_Final(digest.data(), Context.Get()) == 1); + return digest; + } + } + namespace NSha224 { + static_assert(DIGEST_LENGTH == SHA224_DIGEST_LENGTH); + + TDigest Calc(const void* data, size_t dataSize) { + TDigest digest; + Y_ENSURE(SHA224(static_cast(data), dataSize, digest.data()) != nullptr); + return digest; + } + + TCalcer::TCalcer() + : Context{new SHA256state_st} { + Y_ENSURE(SHA224_Init(Context.Get()) == 1); + } + + TCalcer::~TCalcer() { + } + + void TCalcer::Update(const void* data, size_t dataSize) { + Y_ENSURE(SHA224_Update(Context.Get(), data, dataSize) == 1); + } + + TDigest TCalcer::Final() { + TDigest digest; + Y_ENSURE(SHA224_Final(digest.data(), Context.Get()) == 1); + return digest; + } + } +} diff --git a/library/cpp/openssl/crypto/sha.h b/library/cpp/openssl/crypto/sha.h new file mode 100644 index 00000000000..49ec8bd009e --- /dev/null +++ b/library/cpp/openssl/crypto/sha.h @@ -0,0 +1,112 @@ +#pragma once + +#include +#include +#include + +#include + +struct SHAstate_st; +struct SHA256state_st; + +namespace NOpenSsl::NSha1 { + constexpr size_t DIGEST_LENGTH = 20; + using TDigest = std::array; + + // not fragmented input + TDigest Calc(const void* data, size_t dataSize); + + inline TDigest Calc(TStringBuf s) { + return Calc(s.data(), s.length()); + } + + // fragmented input + class TCalcer { + public: + TCalcer(); + ~TCalcer(); + void Update(const void* data, size_t dataSize); + + void Update(TStringBuf s) { + Update(s.data(), s.length()); + } + + template + void UpdateWithPodValue(const T& value) { + Update(&value, sizeof(value)); + } + + TDigest Final(); + + private: + THolder Context; + }; +} + +namespace NOpenSsl::NSha256 { + constexpr size_t DIGEST_LENGTH = 32; + using TDigest = std::array; + + // not fragmented input + TDigest Calc(const void* data, size_t dataSize); + + inline TDigest Calc(TStringBuf s) { + return Calc(s.data(), s.length()); + } + + // fragmented input + class TCalcer { + public: + TCalcer(); + ~TCalcer(); + void Update(const void* data, size_t dataSize); + + void Update(TStringBuf s) { + Update(s.data(), s.length()); + } + + template + void UpdateWithPodValue(const T& value) { + Update(&value, sizeof(value)); + } + + TDigest Final(); + + private: + THolder Context; + }; +} + +namespace NOpenSsl::NSha224 { + constexpr size_t DIGEST_LENGTH = 28; + using TDigest = std::array; + + // not fragmented input + TDigest Calc(const void* data, size_t dataSize); + + inline TDigest Calc(TStringBuf s) { + return Calc(s.data(), s.length()); + } + + // fragmented input + class TCalcer { + public: + TCalcer(); + ~TCalcer(); + void Update(const void* data, size_t dataSize); + + void Update(TStringBuf s) { + Update(s.data(), s.length()); + } + + template + void UpdateWithPodValue(const T& value) { + Update(&value, sizeof(value)); + } + + TDigest Final(); + + private: + THolder Context; + }; +} diff --git a/library/cpp/streams/zstd/zstd.h b/library/cpp/streams/zstd/zstd.h index 667a0494b71..6987bbcacc1 100644 --- a/library/cpp/streams/zstd/zstd.h +++ b/library/cpp/streams/zstd/zstd.h @@ -13,8 +13,8 @@ class TZstdCompress: public IOutputStream { public: /** - @param slave stream to write compressed data to - @param quality, higher quality - slower but better compression. + @param slave stream to write compressed data to + @param quality higher quality - slower but better compression. 0 is default compression (see constant ZSTD_CLEVEL_DEFAULT(3)) max compression is ZSTD_MAX_CLEVEL (22) */ diff --git a/library/cpp/string_utils/quote/quote.cpp b/library/cpp/string_utils/quote/quote.cpp index 91e31da1102..362b5c0bafa 100644 --- a/library/cpp/string_utils/quote/quote.cpp +++ b/library/cpp/string_utils/quote/quote.cpp @@ -4,6 +4,8 @@ #include #include +#include + /* note: (x & 0xdf) makes x upper case */ #define GETXC \ do { \ @@ -181,6 +183,73 @@ static inline It1 Unescape(It1 to, It2 from, It3 end, FromHex fromHex) { return to; } +static constexpr auto x2d = [] { + std::array values{}; + for (auto& value : values) { + value = 0xFF; + } + for (size_t i = 0; i < 10; ++i) { + values['0' + i] = i; + } + for (size_t i = 0; i < 6; ++i) { + values['A' + i] = values['a' + i] = i + 10; + } + return values; +}(); + +static inline void UnescapeHelper(char*& to, const char*& from, const char* end) { + unsigned char c = *from++; + if (c == '%') { + if (end - from >= 2) { + unsigned char hi = x2d[(unsigned char)from[0]]; + unsigned char lo = x2d[(unsigned char)from[1]]; + if ((hi | lo) < 16) { + c = (hi << 4) | lo; + from += 2; + } + } + } else { + c = (c == '+') ? ' ' : c; + } + *to++ = c; +} + +static inline char* Unescape(char* to, const char* from, const char* end, TFromHexLenLimited) { + constexpr size_t BLOCK = 16; + + while (from + BLOCK <= end) { + uint32_t unescape_count = 0; + for (size_t i = 0; i < BLOCK; ++i) { + unescape_count |= (from[i] == '%'); + } + + if (unescape_count == 0) { + unsigned char src[BLOCK]; + for (size_t i = 0; i < BLOCK; ++i) { + src[i] = from[i]; + } + for (size_t i = 0; i < BLOCK; ++i) { + unsigned char c = src[i]; + to[i] = (c == '+') ? ' ' : c; + } + to += BLOCK; + from += BLOCK; + } else { + const char* blockEnd = from + BLOCK; + while (from < blockEnd) { + UnescapeHelper(to, from, end); + } + } + } + + while (from != end) { + UnescapeHelper(to, from, end); + } + + *to = 0; + return to; +} + // CGIEscape returns pointer to the end of the result string // so as it could be possible to populate single long buffer // with several calls to CGIEscape in a row. diff --git a/library/cpp/string_utils/quote/ut/quote_ut.cpp b/library/cpp/string_utils/quote/ut/quote_ut.cpp index b0773ebe996..39da3fb20cd 100644 --- a/library/cpp/string_utils/quote/ut/quote_ut.cpp +++ b/library/cpp/string_utils/quote/ut/quote_ut.cpp @@ -2,6 +2,8 @@ #include +#include + Y_UNIT_TEST_SUITE(TCGIEscapeTest) { Y_UNIT_TEST(ReturnsEndOfTo) { char r[10]; @@ -48,6 +50,91 @@ Y_UNIT_TEST_SUITE(TCGIEscapeTest) { } Y_UNIT_TEST_SUITE(TCGIUnescapeTest) { + Y_UNIT_TEST(BlockBoundaries) { + const TStringBuf fragments[] = {"%", "%3", "%3g", "%3D", "%00", "%ff", "+", "%%41", "%+1"}; + for (size_t offset = 0; offset < 48; ++offset) { + for (const auto fragment : fragments) { + for (size_t tail = 0; tail < 33; ++tail) { + TString input(offset, 'a'); + input += fragment; + input += TString(tail, '+'); + TString expected(input.size() + 1, '\0'); + char* expectedBegin = expected.begin(); + expected.resize(CGIUnescape(expectedBegin, input.c_str()) - expectedBegin); + UNIT_ASSERT_VALUES_EQUAL(CGIUnescapeRet(input), expected); + TVector bounded(input.data(), input.data() + input.size()); + TString result(input.size() + 1, '\0'); + char* resultBegin = result.begin(); + char* resultEnd = CGIUnescape(resultBegin, bounded.data(), bounded.size()); + UNIT_ASSERT_VALUES_EQUAL(TStringBuf(resultBegin, resultEnd), expected); + UNIT_ASSERT_VALUES_EQUAL(*resultEnd, '\0'); + CGIUnescape(input); + UNIT_ASSERT_VALUES_EQUAL(input, expected); + } + } + } + } + + Y_UNIT_TEST(AllHexPairsAtBlockBoundary) { + const auto hexValue = [](unsigned char c) -> int { + if (c >= '0' && c <= '9') { + return c - '0'; + } + if (c >= 'A' && c <= 'F') { + return c - 'A' + 10; + } + if (c >= 'a' && c <= 'f') { + return c - 'a' + 10; + } + return -1; + }; + for (size_t offset : {14, 15}) { + for (unsigned int first = 0; first < 256; ++first) { + for (unsigned int second = 0; second < 256; ++second) { + TString input(offset, 'a'); + input += '%'; + input += static_cast(first); + input += static_cast(second); + input += TString(32, '+'); + + TString expected(offset, 'a'); + const int hi = hexValue(first); + const int lo = hexValue(second); + if (hi >= 0 && lo >= 0) { + expected += static_cast(hi * 16 + lo); + } else { + expected += '%'; + expected += first == '+' ? ' ' : static_cast(first); + expected += second == '+' ? ' ' : static_cast(second); + } + expected += TString(32, ' '); + + UNIT_ASSERT_VALUES_EQUAL(CGIUnescapeRet(input), expected); + CGIUnescape(input); + UNIT_ASSERT_VALUES_EQUAL(input, expected); + } + } + } + } + + Y_UNIT_TEST(AllBytesInBlocks) { + TString input; + for (size_t i = 0; i < 256; ++i) { + input += static_cast(i); + } + TString escaped = CGIEscapeRet(input); + UNIT_ASSERT_VALUES_EQUAL(CGIUnescapeRet(escaped), input); + CGIUnescape(escaped); + UNIT_ASSERT_VALUES_EQUAL(escaped, input); + + // Exercise block copying with embedded zeros and high bytes as well. + TString expected = input; + expected[static_cast('+')] = ' '; + UNIT_ASSERT_VALUES_EQUAL(CGIUnescapeRet(input), expected); + CGIUnescape(input); + UNIT_ASSERT_VALUES_EQUAL(input, expected); + } + Y_UNIT_TEST(StringBuf) { char tmp[100]; diff --git a/library/cpp/string_utils/url/url.cpp b/library/cpp/string_utils/url/url.cpp index 7d7d7dbd90b..e3a7c36b913 100644 --- a/library/cpp/string_utils/url/url.cpp +++ b/library/cpp/string_utils/url/url.cpp @@ -175,6 +175,22 @@ static inline TStringBuf GetHostAndPortImpl(const TStringBuf url) { }; const auto& nonHostCharacters = *Singleton(); + + // Handle IPv6 addresses enclosed in brackets (RFC 3986): the host is [...], + // optionally followed by :port. Colons inside brackets are not delimiters. + if (!urlNoScheme.empty() && urlNoScheme[0] == '[') { + const size_t closeBracket = urlNoScheme.find(']'); + if (closeBracket != TStringBuf::npos) { + if constexpr (KeepPort) { + // Include port after ]: find end at first /;?# + const char* end = nonHostCharacters.brk(urlNoScheme.begin() + closeBracket + 1, urlNoScheme.end()); + return urlNoScheme.Head(end - urlNoScheme.data()); + } else { + return urlNoScheme.Head(closeBracket + 1); + } + } + } + const char* firstNonHostCharacter = nonHostCharacters.brk(urlNoScheme.begin(), urlNoScheme.end()); if (firstNonHostCharacter != urlNoScheme.end()) { @@ -216,7 +232,7 @@ TStringBuf GetSchemeHostAndPort(const TStringBuf url Y_LIFETIME_BOUND, bool trim TStringBuf hostAndPort = GetHostAndPort(url.Tail(schemeSize)); if (trimDefaultPort) { - const size_t pos = hostAndPort.find(':'); + const size_t pos = hostAndPort.rfind(':'); if (pos != TStringBuf::npos) { const bool isHttps = (scheme == TStringBuf("https://")); diff --git a/library/cpp/string_utils/url/url.h b/library/cpp/string_utils/url/url.h index 5e707756c3b..a79f311ee88 100644 --- a/library/cpp/string_utils/url/url.h +++ b/library/cpp/string_utils/url/url.h @@ -5,19 +5,16 @@ namespace NUrl { - /** - * Splits URL to host and path - * Example: - * auto [host, path] = SplitUrlToHostAndPath(url); - * - * @param[in] url any URL - * @param[out] parsed host and path - */ struct TSplitUrlToHostAndPathResult { TStringBuf host; TStringBuf path; }; + /** + * Splits URL to host and path + * Example: + * auto [host, path] = SplitUrlToHostAndPath(url); + */ Y_PURE_FUNCTION TSplitUrlToHostAndPathResult SplitUrlToHostAndPath(const TStringBuf url Y_LIFETIME_BOUND); @@ -86,9 +83,10 @@ TStringBuf GetSchemeHostAndPort(const TStringBuf url Y_LIFETIME_BOUND, bool trim * @param[in] url any URL * @param[out] host parsed host * @param[out] path parsed path - */ + * @{ */ void SplitUrlToHostAndPath(const TStringBuf url, TStringBuf& host, TStringBuf& path); void SplitUrlToHostAndPath(const TStringBuf url, TString& host, TString& path); +/** @} */ /** * Separates URL into url prefix, query (aka cgi params list), and fragment (aka part after #) diff --git a/library/cpp/string_utils/url/url_ut.cpp b/library/cpp/string_utils/url/url_ut.cpp index 7980a36e995..23ca91b257a 100644 --- a/library/cpp/string_utils/url/url_ut.cpp +++ b/library/cpp/string_utils/url/url_ut.cpp @@ -65,6 +65,23 @@ Y_UNIT_TEST_SUITE(TUtilUrlTest) { UNIT_ASSERT_VALUES_EQUAL("http://ya.ru:81", GetSchemeHostAndPort("http://ya.ru:81", /*trimHttp*/false)); UNIT_ASSERT_VALUES_EQUAL("http://ya.ru:81", GetSchemeHostAndPort("http://ya.ru:81", /*trimHttp*/false, /*trimDefaultPort*/false)); + // trimDefaultPort=true with IPv4 + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetSchemeHostAndPort("http://192.168.1.1:80/bebe")); + UNIT_ASSERT_VALUES_EQUAL("https://192.168.1.1", GetSchemeHostAndPort("https://192.168.1.1:443/bebe")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1:8080", GetSchemeHostAndPort("http://192.168.1.1:8080/bebe")); + UNIT_ASSERT_VALUES_EQUAL("https://192.168.1.1:8443", GetSchemeHostAndPort("https://192.168.1.1:8443/bebe")); + UNIT_ASSERT_VALUES_EQUAL("http://192.168.1.1", GetSchemeHostAndPort("http://192.168.1.1:80/bebe", /*trimHttp*/false)); + UNIT_ASSERT_VALUES_EQUAL("http://192.168.1.1:8080", GetSchemeHostAndPort("http://192.168.1.1:8080/bebe", /*trimHttp*/false)); + // trimDefaultPort=true with IPv6 (port is after the closing bracket) + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetSchemeHostAndPort("http://[::1]:80/bebe")); + UNIT_ASSERT_VALUES_EQUAL("https://[::1]", GetSchemeHostAndPort("https://[::1]:443/bebe")); + UNIT_ASSERT_VALUES_EQUAL("[::1]:8080", GetSchemeHostAndPort("http://[::1]:8080/bebe")); + UNIT_ASSERT_VALUES_EQUAL("https://[::1]:8443", GetSchemeHostAndPort("https://[::1]:8443/bebe")); + UNIT_ASSERT_VALUES_EQUAL("http://[::1]", GetSchemeHostAndPort("http://[::1]:80/bebe", /*trimHttp*/false)); + UNIT_ASSERT_VALUES_EQUAL("http://[::1]:8080", GetSchemeHostAndPort("http://[::1]:8080/bebe", /*trimHttp*/false)); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetSchemeHostAndPort("http://[2001:db8::1]:80/bebe")); + UNIT_ASSERT_VALUES_EQUAL("https://[2001:db8::1]", GetSchemeHostAndPort("https://[2001:db8::1]:443/bebe")); + // irl RFC3986 sometimes gets ignored UNIT_ASSERT_VALUES_EQUAL("pravda-kmv.ru", GetSchemeHostAndPort("pravda-kmv.ru?page=news&id=6973")); UNIT_ASSERT_VALUES_EQUAL("pravda-kmv.ru", GetSchemeHostAndPort("pravda-kmv.ru?page=news&id=6973", /*trimHttp*/false)); @@ -365,4 +382,105 @@ Y_UNIT_TEST_SUITE(TUtilUrlTest) { UNIT_ASSERT_VALUES_EQUAL(false, DoesUrlPathStartWithToken("http://bebe", "bebe")); UNIT_ASSERT_VALUES_EQUAL(false, DoesUrlPathStartWithToken("https://bebe/", "bebe")); } + + Y_UNIT_TEST(TestGetHostIp) { + // IPv4 + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetHost("192.168.1.1/path")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetHost("192.168.1.1:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetHost("http://192.168.1.1/path")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetHost("https://192.168.1.1:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetHost("192.168.1.1")); + // IPv6 (RFC 3986: address is enclosed in brackets, port is separated by colon after brackets) + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetHost("[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetHost("[::1]:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetHost("http://[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetHost("https://[::1]:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetHost("[::1]")); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetHost("http://[2001:db8::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetHost("http://[2001:db8::1]:8080/path")); + } + + Y_UNIT_TEST(TestGetSchemeHostIp) { + // IPv4 + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetSchemeHost("http://192.168.1.1/path")); + UNIT_ASSERT_VALUES_EQUAL("http://192.168.1.1", GetSchemeHost("http://192.168.1.1/path", false)); + UNIT_ASSERT_VALUES_EQUAL("https://192.168.1.1", GetSchemeHost("https://192.168.1.1/path")); + UNIT_ASSERT_VALUES_EQUAL("https://192.168.1.1", GetSchemeHost("https://192.168.1.1:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("192.168.1.1", GetSchemeHost("192.168.1.1/path")); + // IPv6 + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetSchemeHost("http://[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("http://[::1]", GetSchemeHost("http://[::1]/path", false)); + UNIT_ASSERT_VALUES_EQUAL("https://[::1]", GetSchemeHost("https://[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("https://[::1]", GetSchemeHost("https://[::1]:8080/path")); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetSchemeHost("[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("https://[2001:db8::1]", GetSchemeHost("https://[2001:db8::1]/path")); + } + + Y_UNIT_TEST(TestGetPathAndQueryIp) { + // IPv4 + UNIT_ASSERT_VALUES_EQUAL("/path", GetPathAndQuery("192.168.1.1/path")); + UNIT_ASSERT_VALUES_EQUAL("/path?query", GetPathAndQuery("http://192.168.1.1/path?query")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("192.168.1.1")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("http://192.168.1.1")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("192.168.1.1:8080")); + // IPv6 + UNIT_ASSERT_VALUES_EQUAL("/path", GetPathAndQuery("[::1]/path")); + UNIT_ASSERT_VALUES_EQUAL("/path?query", GetPathAndQuery("http://[::1]/path?query")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("[::1]")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("http://[::1]")); + UNIT_ASSERT_VALUES_EQUAL("/", GetPathAndQuery("[::1]:8080")); + UNIT_ASSERT_VALUES_EQUAL("/path", GetPathAndQuery("http://[2001:db8::1]/path")); + } + + Y_UNIT_TEST(TestGetDomainIp) { + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetDomain("[::1]")); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetDomain("[2001:db8::1]")); + } + + Y_UNIT_TEST(TestGetParentDomainIp) { + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetParentDomain("[::1]", 1)); + UNIT_ASSERT_VALUES_EQUAL("[::1]", GetParentDomain("[::1]", 2)); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetParentDomain("[2001:db8::1]", 1)); + UNIT_ASSERT_VALUES_EQUAL("[2001:db8::1]", GetParentDomain("[2001:db8::1]", 2)); + } + + Y_UNIT_TEST(TestSeparateUrlFromQueryAndFragmentIp) { + // IPv4 + { + TStringBuf sanitizedUrl, query, fragment; + SeparateUrlFromQueryAndFragment("http://192.168.1.1/path?param=val#frag", sanitizedUrl, query, fragment); + UNIT_ASSERT_STRINGS_EQUAL(sanitizedUrl, "http://192.168.1.1/path"); + UNIT_ASSERT_STRINGS_EQUAL(query, "param=val"); + UNIT_ASSERT_STRINGS_EQUAL(fragment, "frag"); + } + { + TStringBuf sanitizedUrl, query, fragment; + SeparateUrlFromQueryAndFragment("192.168.1.1/path?param=val", sanitizedUrl, query, fragment); + UNIT_ASSERT_STRINGS_EQUAL(sanitizedUrl, "192.168.1.1/path"); + UNIT_ASSERT_STRINGS_EQUAL(query, "param=val"); + UNIT_ASSERT_STRINGS_EQUAL(fragment, ""); + } + // IPv6 + { + TStringBuf sanitizedUrl, query, fragment; + SeparateUrlFromQueryAndFragment("http://[::1]/path?param=val#frag", sanitizedUrl, query, fragment); + UNIT_ASSERT_STRINGS_EQUAL(sanitizedUrl, "http://[::1]/path"); + UNIT_ASSERT_STRINGS_EQUAL(query, "param=val"); + UNIT_ASSERT_STRINGS_EQUAL(fragment, "frag"); + } + { + TStringBuf sanitizedUrl, query, fragment; + SeparateUrlFromQueryAndFragment("[::1]/path?param=val", sanitizedUrl, query, fragment); + UNIT_ASSERT_STRINGS_EQUAL(sanitizedUrl, "[::1]/path"); + UNIT_ASSERT_STRINGS_EQUAL(query, "param=val"); + UNIT_ASSERT_STRINGS_EQUAL(fragment, ""); + } + { + TStringBuf sanitizedUrl, query, fragment; + SeparateUrlFromQueryAndFragment("http://[2001:db8::1]/path#frag", sanitizedUrl, query, fragment); + UNIT_ASSERT_STRINGS_EQUAL(sanitizedUrl, "http://[2001:db8::1]/path"); + UNIT_ASSERT_STRINGS_EQUAL(query, ""); + UNIT_ASSERT_STRINGS_EQUAL(fragment, "frag"); + } + } } diff --git a/library/cpp/threading/chunk_queue/README.md b/library/cpp/threading/chunk_queue/README.md index 62a414e4f3e..b842f5dc569 100644 --- a/library/cpp/threading/chunk_queue/README.md +++ b/library/cpp/threading/chunk_queue/README.md @@ -43,17 +43,13 @@ The cheapest MPSC variant: elements are spread across internal partitions, so FI The same as `TRelaxedManyOneQueue`, but with multiple readers allowed. -### Pointer queues: `TAutoOneOneQueue` and friends - -[`TAutoQueueBase`](queue.h) is a wrapper for queues of owning pointers: `Enqueue(TAutoPtr)` takes ownership, `Dequeue` returns a `TAutoPtr`, and the queue destructor deletes all remaining elements. Aliases: `TAutoOneOneQueue`, `TAutoManyOneQueue`, `TAutoManyManyQueue`, `TAutoRelaxedManyOneQueue`, `TAutoRelaxedManyManyQueue`. - ## Example See usage examples: [`queue_ut.cpp`](queue_ut.cpp). ## Template parameters -- `T` — the element type (stored by value; for pointer ownership use the `TAuto*` aliases); +- `T` — the element type (stored by value); - `ChunkSize` — chunk size in bytes (default 4 KiB); a larger chunk means fewer allocations but more overhead for small queues; - `Concurrency` — the number of internal partitions in the `TMany*`/`TRelaxed*` queues (default 4); scale it with the number of writers; - `TLock` (`TManyManyQueue` only) — the lock type (default `TAdaptiveLock`). diff --git a/library/cpp/threading/chunk_queue/queue.h b/library/cpp/threading/chunk_queue/queue.h index c896e85361f..16d0e77995a 100644 --- a/library/cpp/threading/chunk_queue/queue.h +++ b/library/cpp/threading/chunk_queue/queue.h @@ -1,7 +1,6 @@ #pragma once #include -#include #include #include #include @@ -11,6 +10,7 @@ #include #include +#include #if defined(_MSC_VER) && !defined(__clang__) #include @@ -63,13 +63,11 @@ namespace NThreading { static_assert(std::atomic::is_always_lock_free, "TAtomicRef requires a lock-free std::atomic"); static_assert( !std::is_const::value && !std::is_volatile::value, - "TAtomicRef cannot be used with cv-qualified types" - ); + "TAtomicRef cannot be used with cv-qualified types"); static_assert(alignof(TT) == sizeof(TT), "TT must have natural alignment"); static_assert( (sizeof(TT) == 1) || (sizeof(TT) == 2) || (sizeof(TT) == 4) || (sizeof(TT) == 8), - "sizeof(TT) different from 1, 2, 4 or 8 is not supported" - ); + "sizeof(TT) different from 1, 2, 4 or 8 is not supported"); #if !defined(_MSC_VER) || defined(__clang__) static_assert(static_cast(std::memory_order_release) == __ATOMIC_RELEASE); static_assert(static_cast(std::memory_order_acquire) == __ATOMIC_ACQUIRE); @@ -87,8 +85,7 @@ namespace NThreading { #if defined(_MSC_VER) && !defined(__clang__) Y_ABORT_IF( order == std::memory_order_acq_rel || order == std::memory_order_release, - "load: Invalid memory order" - ); + "load: Invalid memory order"); if (order == std::memory_order_seq_cst) { if constexpr (sizeof(TT) == 1) { @@ -133,11 +130,8 @@ namespace NThreading { void store(TT desired, std::memory_order order) noexcept { #if defined(_MSC_VER) && !defined(__clang__) Y_ABORT_IF( - order == std::memory_order_acq_rel - || order == std::memory_order_consume - || order == std::memory_order_acquire, - "store: Invalid memory order" - ); + order == std::memory_order_acq_rel || order == std::memory_order_consume || order == std::memory_order_acquire, + "store: Invalid memory order"); if (order == std::memory_order_seq_cst) { if constexpr (sizeof(TT) == 1) { @@ -646,56 +640,4 @@ namespace NThreading { return false; } }; - - //////////////////////////////////////////////////////////////////////////////// - // Simple wrapper to deal with AutoPtrs - - template - class TAutoQueueBase: private TNonCopyable { - private: - TImpl Impl; - - public: - using TItem = TAutoPtr; - - ~TAutoQueueBase() { - TItem value; - while (Dequeue(value)) { - // do nothing - } - } - - void Enqueue(TItem value) { - Impl.Enqueue(value.Get()); - Y_UNUSED(value.Release()); - } - - bool Dequeue(TItem& value) { - T* ptr = nullptr; - if (Impl.Dequeue(ptr)) { - value.Reset(ptr); - return true; - } - return false; - } - - bool IsEmpty() { - return Impl.IsEmpty(); - } - }; - - template - using TAutoOneOneQueue = TAutoQueueBase>; - - template - using TAutoManyOneQueue = TAutoQueueBase>; - - template - using TAutoManyManyQueue = TAutoQueueBase>; - - template - using TAutoRelaxedManyOneQueue = TAutoQueueBase>; - - template - using TAutoRelaxedManyManyQueue = TAutoQueueBase>; } // namespace NThreading diff --git a/library/cpp/threading/future/legacy_future.h b/library/cpp/threading/future/legacy_future.h index 6f1eabad73b..18cf39553b1 100644 --- a/library/cpp/threading/future/legacy_future.h +++ b/library/cpp/threading/future/legacy_future.h @@ -41,7 +41,7 @@ namespace NThreading { inline void Join() { if (Thread_) { Thread_->Join(); - Thread_.Destroy(); + Thread_.reset(); } } diff --git a/library/cpp/threading/future/subscription/README.md b/library/cpp/threading/future/subscription/README.md index 6f547926854..062bbb4ca4e 100644 --- a/library/cpp/threading/future/subscription/README.md +++ b/library/cpp/threading/future/subscription/README.md @@ -40,7 +40,7 @@ Wait privimitives could be constructed either from an initializer_list or from a Subscriptions manager --------------------- -The subscription manager can manage multiple links beetween futures and callbacks. Multiple managed subscriptions to a single future shares just a single underlying subscription to the future. That allows dynamic creation and deletion of subscriptions and efficient implementation of different wait primitives. +The subscription manager can manage multiple links between futures and callbacks. Multiple managed subscriptions to a single future shares just a single underlying subscription to the future. That allows dynamic creation and deletion of subscriptions and efficient implementation of different wait primitives. The subscription manager could be used in the following way: 1. Subscribe to a single future: diff --git a/library/cpp/threading/future/subscription/wait.h b/library/cpp/threading/future/subscription/wait.h index 533bab9d8d9..1e9ad3793ca 100644 --- a/library/cpp/threading/future/subscription/wait.h +++ b/library/cpp/threading/future/subscription/wait.h @@ -1,4 +1,4 @@ -#pragma once +#pragma once #include "subscription.h" @@ -56,7 +56,7 @@ class TWait : public TThrRefBase { private: //! Performs a subscription to the given futures /** Lock should not be acquired! - @param future - The futures to subscribe to + @param futures - The futures to subscribe to @param callback - The callback to call for each future **/ template diff --git a/library/cpp/yt/compact_containers/compact_flat_map.h b/library/cpp/yt/compact_containers/compact_flat_map.h index b6c8adad93c..fc3e265cc9a 100644 --- a/library/cpp/yt/compact_containers/compact_flat_map.h +++ b/library/cpp/yt/compact_containers/compact_flat_map.h @@ -1,5 +1,6 @@ #pragma once +#include "comparison.h" #include "compact_vector.h" #include @@ -8,20 +9,6 @@ namespace NYT { /////////////////////////////////////////////////////////////////////////////// -namespace NDetail { - -template -concept CHasIsTransparentFlag = requires { - typename T::is_transparent; -}; - -template -concept CComparisonAllowed = std::same_as || CHasIsTransparentFlag; - -} // namespace NDetail - -/////////////////////////////////////////////////////////////////////////////// - //! A flat map implementation over TCompactVector that tries to keep data inline. /*! * Similarly to SmallSet, this is implemented via binary search over a sorted diff --git a/library/cpp/yt/compact_containers/compact_map-inl.h b/library/cpp/yt/compact_containers/compact_map-inl.h new file mode 100644 index 00000000000..fc449fa09d4 --- /dev/null +++ b/library/cpp/yt/compact_containers/compact_map-inl.h @@ -0,0 +1,911 @@ +#ifndef COMPACT_MAP_INL_H_ +#error "Direct inclusion of this file is not allowed, include compact_map.h" +// For the sake of sane code completion. +#include "compact_map.h" +#endif + +namespace NYT { + +//////////////////////////////////////////////////////////////////////////////// + +template +template +class TCompactMap::TIteratorBase +{ +private: + friend class TCompactMap; + friend class TIteratorBase; + friend class iterator; + friend class const_iterator; + + using TVecIter = std::conditional_t; + using TMapIter = std::conditional_t; + using TValueType = typename TCompactMap::value_type; + using TReference = std::conditional_t; + using TPointer = std::conditional_t; + + union + { + TVecIter VIter; + TMapIter MIter; + }; + bool Small_; + + TIteratorBase() + : VIter() + , Small_(true) + { } + + explicit TIteratorBase(TVecIter it) + : VIter(it) + , Small_(true) + { } + + explicit TIteratorBase(TMapIter it) + : MIter(it) + , Small_(false) + { } + +public: + using difference_type = std::ptrdiff_t; + using value_type = TValueType; + using reference = TReference; + using pointer = TPointer; + using iterator_category = std::bidirectional_iterator_tag; + + reference operator*() const + { + if (Small_) { + return *VIter; + } + return *MIter; + } + + pointer operator->() const + { + if (Small_) { + return &*VIter; + } + return &*MIter; + } + + TIteratorBase& operator++() + { + if (Small_) { + ++VIter; + } else { + ++MIter; + } + return *this; + } + + TIteratorBase operator++(int) + { + TIteratorBase tmp = *this; + ++*this; + return tmp; + } + + TIteratorBase& operator--() + { + if (Small_) { + --VIter; + } else { + --MIter; + } + return *this; + } + + TIteratorBase operator--(int) + { + TIteratorBase tmp = *this; + --*this; + return tmp; + } + + template + bool operator==(const TIteratorBase& other) const + { + if (Small_ != other.Small_) { + return false; + } + return Small_ ? VIter == other.VIter : MIter == other.MIter; + } +}; + +template +class TCompactMap::iterator + : public TCompactMap::template TIteratorBase +{ +public: + iterator() = default; + + iterator& operator++() + { + TBase::operator++(); + return *this; + } + + iterator operator++(int) + { + auto result = *this; + ++*this; + return result; + } + + iterator& operator--() + { + TBase::operator--(); + return *this; + } + + iterator operator--(int) + { + auto result = *this; + --*this; + return result; + } + +private: + friend class TCompactMap; + using TBase = typename TCompactMap::template TIteratorBase; + + explicit iterator(TArrayIterator it) + : TBase(it) + { } + + explicit iterator(TMapIterator it) + : TBase(it) + { } +}; + +template +class TCompactMap::const_iterator + : public TCompactMap::template TIteratorBase +{ +public: + const_iterator() = default; + + const_iterator& operator++() + { + TBase::operator++(); + return *this; + } + + const_iterator operator++(int) + { + auto result = *this; + ++*this; + return result; + } + + const_iterator& operator--() + { + TBase::operator--(); + return *this; + } + + const_iterator operator--(int) + { + auto result = *this; + --*this; + return result; + } + + // Allow implicit conversion from iterator to const_iterator. + const_iterator(const iterator& other) + : TBase(other.Small_ ? TBase(other.VIter) : TBase(other.MIter)) + { } + +private: + friend class TCompactMap; + using TBase = typename TCompactMap::template TIteratorBase; + + explicit const_iterator(TArrayConstIterator it) + : TBase(it) + { } + + explicit const_iterator(TMapConstIterator it) + : TBase(it) + { } +}; + +//////////////////////////////////////////////////////////////////////////////// + +template +auto TCompactMap::ArrayData() -> TArrayIterator +{ + static_assert(sizeof(TStorage) == sizeof(TArrayValue), "Inline storage size mismatch"); + static_assert(alignof(TStorage) == alignof(TArrayValue), "Inline storage alignment mismatch"); + return reinterpret_cast(ArrayStorage_.data()); +} + +template +auto TCompactMap::ArrayData() const -> TArrayConstIterator +{ + return reinterpret_cast(ArrayStorage_.data()); +} + +template +auto TCompactMap::ArrayBegin() -> TArrayIterator +{ + return ArrayData(); +} + +template +auto TCompactMap::ArrayBegin() const -> TArrayConstIterator +{ + return ArrayData(); +} + +template +auto TCompactMap::ArrayEnd() -> TArrayIterator +{ + return ArrayData() + ArraySize_; +} + +template +auto TCompactMap::ArrayEnd() const -> TArrayConstIterator +{ + return ArrayData() + ArraySize_; +} + +template +void TCompactMap::ClearArray() +{ + auto* data = ArrayData(); + for (size_type index = 0; index < ArraySize_; ++index) { + std::destroy_at(data + index); + } + ArraySize_ = 0; +} + +template +void TCompactMap::CopyArrayFrom(const TCompactMap& other) +{ + auto* data = ArrayData(); + auto* otherData = other.ArrayData(); + + try { + for (size_type index = 0; index < other.ArraySize_; ++index) { + std::construct_at(data + index, otherData[index]); + ++ArraySize_; + } + } catch (...) { + ClearArray(); + throw; + } +} + +template +void TCompactMap::MoveArrayFrom(TCompactMap&& other) +{ + auto* data = ArrayData(); + auto* otherData = other.ArrayData(); + + try { + for (size_type index = 0; index < other.ArraySize_; ++index) { + std::construct_at(data + index, std::move(otherData[index])); + ++ArraySize_; + } + } catch (...) { + ClearArray(); + throw; + } + + other.ClearArray(); +} + +//////////////////////////////////////////////////////////////////////////////// + +template +TCompactMap::TCompactMap(const TCompactMap& other) + : Compare_(other.Compare_) +{ + if (other.IsSmall()) { + CopyArrayFrom(other); + } else { + Map_ = other.Map_; + MapMode_ = true; + } +} + +template +TCompactMap::TCompactMap(TCompactMap&& other) noexcept( + std::is_nothrow_move_constructible_v && + std::is_nothrow_move_constructible_v) + : Map_(std::move(other.Map_)) + , Compare_(std::move(other.Compare_)) +{ + if (other.IsSmall()) { + MoveArrayFrom(std::move(other)); + } else { + MapMode_ = true; + + other.Map_.clear(); + other.MapMode_ = false; + } +} + +template +TCompactMap& +TCompactMap::operator=(const TCompactMap& other) +{ + if (this == &other) { + return *this; + } + + ClearArray(); + Map_.clear(); + MapMode_ = false; + + Compare_ = other.Compare_; + + if (other.IsSmall()) { + CopyArrayFrom(other); + } else { + Map_ = other.Map_; + MapMode_ = true; + } + + return *this; +} + +template +TCompactMap& +TCompactMap::operator=(TCompactMap&& other) noexcept( + std::is_nothrow_move_constructible_v && + std::is_nothrow_move_assignable_v) +{ + if (this == &other) { + return *this; + } + + ClearArray(); + Map_.clear(); + MapMode_ = false; + + Compare_ = std::move(other.Compare_); + + if (other.IsSmall()) { + MoveArrayFrom(std::move(other)); + } else { + Map_ = std::move(other.Map_); + MapMode_ = true; + + other.Map_.clear(); + other.MapMode_ = false; + } + + return *this; +} + +template +TCompactMap::~TCompactMap() +{ + ClearArray(); +} + +template +template +TCompactMap::TCompactMap(TIt first, TIt last) +{ + insert(first, last); +} + +template +TCompactMap::TCompactMap(std::initializer_list init) +{ + insert(init.begin(), init.end()); +} + +template +bool TCompactMap::empty() const +{ + return IsSmall() ? ArraySize_ == 0 : Map_.empty(); +} + +template +typename TCompactMap::size_type +TCompactMap::size() const +{ + return IsSmall() ? ArraySize_ : Map_.size(); +} + +template +void TCompactMap::clear() +{ + ClearArray(); + Map_.clear(); + MapMode_ = false; +} + +template +auto TCompactMap::begin() -> iterator +{ + return BeginImpl(this); +} + +template +auto TCompactMap::begin() const -> const_iterator +{ + return BeginImpl(this); +} + +template +auto TCompactMap::cbegin() const -> const_iterator +{ + return begin(); +} + +template +auto TCompactMap::end() -> iterator +{ + return EndImpl(this); +} + +template +auto TCompactMap::end() const -> const_iterator +{ + return EndImpl(this); +} + +template +auto TCompactMap::cend() const -> const_iterator +{ + return end(); +} + +template +bool TCompactMap::IsSmall() const +{ + return !MapMode_; +} + +template +void TCompactMap::UpgradeToMap() +{ + if (!IsSmall()) { + return; + } + + TMap newMap(Compare_); + auto* data = ArrayData(); + for (size_type index = 0; index < ArraySize_; ++index) { + newMap.emplace_hint(newMap.end(), data[index].first, std::move(data[index].second)); + } + + ClearArray(); + Map_ = std::move(newMap); + MapMode_ = true; +} + +template +template +std::pair::TArrayConstIterator, bool> +TCompactMap::ArrayLowerBound(const TOtherKey& key) const +{ + return ArrayLowerBoundImpl(ArrayBegin(), ArrayEnd(), key, Compare_); +} + +template +template +std::pair::TArrayIterator, bool> +TCompactMap::ArrayLowerBound(const TOtherKey& key) +{ + return ArrayLowerBoundImpl(ArrayBegin(), ArrayEnd(), key, Compare_); +} + +template +template +auto TCompactMap::EmplaceSmall(TArrayConstIterator pos, TArgs&&... args) -> TArrayIterator +{ + auto index = static_cast(pos - ArrayBegin()); + auto* data = ArrayData(); + + // Simple path: emplace at the end. + if (index == ArraySize_) { + std::construct_at(data + ArraySize_, std::forward(args)...); + ++ArraySize_; + return data + index; + } + + // Construct the new value before modifying the array. + TArrayValue value(std::forward(args)...); + + // Shift elements to make space at index. If relocation throws, destroy the + // shifted suffix and keep the remaining live prefix valid. + std::construct_at(data + ArraySize_, std::move(data[ArraySize_ - 1])); + for (size_type i = ArraySize_ - 1; i > index; --i) { + std::destroy_at(data + i); + try { + std::construct_at(data + i, std::move(data[i - 1])); + } catch (...) { + for (size_type j = i + 1; j <= ArraySize_; ++j) { + std::destroy_at(data + j); + } + ArraySize_ = i; + throw; + } + } + + std::destroy_at(data + index); + try { + std::construct_at(data + index, std::move(value)); + } catch (...) { + for (size_type i = index + 1; i <= ArraySize_; ++i) { + std::destroy_at(data + i); + } + ArraySize_ = index; + throw; + } + + ++ArraySize_; + return data + index; +} + +template +auto TCompactMap::EraseSmall(TArrayConstIterator pos) -> TArrayIterator +{ + return EraseSmall(pos, pos + 1); +} + +template +auto TCompactMap::EraseSmall(TArrayConstIterator first, TArrayConstIterator last) -> TArrayIterator +{ + auto start = static_cast(first - ArrayBegin()); + auto finish = static_cast(last - ArrayBegin()); + auto* data = ArrayData(); + + if (start == finish) { + return data + start; + } + + // Destroy elements in [start, finish). + for (size_type index = start; index < finish; ++index) { + std::destroy_at(data + index); + } + + // Shift elements from [finish, ArraySize_) to [start, ...). + size_type index = finish; + try { + for (; index < ArraySize_; ++index) { + std::construct_at(data + start + (index - finish), std::move(data[index])); + std::destroy_at(data + index); + } + } catch (...) { + // Relocation failed; drop the tail that has not been moved yet. + for (size_type tailIndex = index; tailIndex < ArraySize_; ++tailIndex) { + std::destroy_at(data + tailIndex); + } + ArraySize_ = start + (index - finish); + throw; + } + + ArraySize_ -= finish - start; + return data + start; +} + +template +template +std::pair +TCompactMap::ArrayLowerBoundImpl( + TArrayIteratorType begin, + TArrayIteratorType end, + const TOtherKey& key, + const key_compare& compare) +{ + auto comp = [&compare] (const auto& pair, const auto& k) { + return compare(pair.first, k); + }; + + auto it = std::lower_bound(begin, end, key, comp); + bool found = it != end && !compare(key, it->first); + return {it, found}; +} + +template +template +typename TCompactMap::mapped_type& +TCompactMap::Subscript(TKeyParam&& key) +{ + return TryEmplaceImpl(std::forward(key)).first->second; +} + +template +template +auto TCompactMap::BeginImpl(TSelf* self) + -> std::conditional_t, const_iterator, iterator> +{ + using TResult = std::conditional_t, const_iterator, iterator>; + if (self->IsSmall()) { + return TResult(self->ArrayBegin()); + } + return TResult(self->Map_.begin()); +} + +template +template +auto TCompactMap::EndImpl(TSelf* self) + -> std::conditional_t, const_iterator, iterator> +{ + using TResult = std::conditional_t, const_iterator, iterator>; + if (self->IsSmall()) { + return TResult(self->ArrayEnd()); + } + return TResult(self->Map_.end()); +} + +template +template +auto TCompactMap::FindImpl(TSelf* self, const TOtherKey& key) + -> std::conditional_t, const_iterator, iterator> +{ + using TResult = std::conditional_t, const_iterator, iterator>; + if (self->IsSmall()) { + auto [it, found] = self->ArrayLowerBound(key); + return found ? TResult(it) : TResult(self->ArrayEnd()); + } + return TResult(self->Map_.find(key)); +} + +template +template +auto TCompactMap::AtImpl(TSelf* self, const key_type& key) + -> std::conditional_t, const mapped_type&, mapped_type&> +{ + auto it = FindImpl(self, key); + if (it == EndImpl(self)) { + throw std::out_of_range("TCompactMap::at"); + } + return it->second; +} + +template +template TOtherKey> +typename TCompactMap::size_type +TCompactMap::count(const TOtherKey& key) const +{ + return FindImpl(this, key) == EndImpl(this) ? 0 : 1; +} + +template +template TOtherKey> +bool TCompactMap::contains(const TOtherKey& key) const +{ + return FindImpl(this, key) != EndImpl(this); +} + +template +template TOtherKey> +auto TCompactMap::find(const TOtherKey& key) -> iterator +{ + return FindImpl(this, key); +} + +template +template TOtherKey> +auto TCompactMap::find(const TOtherKey& key) const -> const_iterator +{ + return FindImpl(this, key); +} + +template +std::pair::iterator, bool> +TCompactMap::insert(const value_type& value) +{ + return DoInsert(value); +} + +template +std::pair::iterator, bool> +TCompactMap::insert(value_type&& value) +{ + return DoInsert(std::move(value)); +} + +template +template +std::pair::iterator, bool> +TCompactMap::DoInsert(TValueParam&& value) +{ + if (IsSmall()) { + auto [arrayIt, found] = ArrayLowerBound(value.first); + + if (found) { + return {iterator(arrayIt), false}; + } + + if (ArraySize_ < N) { + auto inserted = EmplaceSmall(arrayIt, std::forward(value)); + return {iterator(inserted), true}; + } + + UpgradeToMap(); + } + + auto [mapIt, inserted] = Map_.insert(std::forward(value)); + return {iterator(mapIt), inserted}; +} + +template +template +std::pair::iterator, bool> +TCompactMap::emplace(TArgs&&... args) +{ + if (IsSmall()) { + TArrayValue value(std::forward(args)...); + auto [arrayIt, found] = ArrayLowerBound(value.first); + + if (found) { + return {iterator(arrayIt), false}; + } + + if (ArraySize_ < N) { + auto inserted = EmplaceSmall(arrayIt, std::move(value)); + return {iterator(inserted), true}; + } + + UpgradeToMap(); + + // We already constructed the value above, so we cannot forward args again below. + auto [mapIt, inserted] = Map_.emplace(std::move(value)); + return {iterator(mapIt), inserted}; + } + + auto [mapIt, inserted] = Map_.emplace(std::forward(args)...); + return {iterator(mapIt), inserted}; +} + +template +template +std::pair::iterator, bool> +TCompactMap::try_emplace(const key_type& key, TArgs&&... args) +{ + return TryEmplaceImpl(key, std::forward(args)...); +} + +template +template +std::pair::iterator, bool> +TCompactMap::try_emplace(key_type&& key, TArgs&&... args) +{ + return TryEmplaceImpl(std::move(key), std::forward(args)...); +} + +template +template +void TCompactMap::insert(TIt first, TIt last) +{ + for (; first != last; ++first) { + insert(*first); + } +} + +template +template +std::pair::iterator, bool> +TCompactMap::insert_or_assign(const key_type& key, TMapped&& obj) +{ + if (IsSmall()) { + auto [arrayIt, found] = ArrayLowerBound(key); + + if (found) { + arrayIt->second = std::forward(obj); + return {iterator(arrayIt), false}; + } + + if (ArraySize_ < N) { + auto inserted = EmplaceSmall(arrayIt, key, std::forward(obj)); + return {iterator(inserted), true}; + } + + UpgradeToMap(); + } + + auto [mapIt, inserted] = Map_.insert_or_assign(key, std::forward(obj)); + return {iterator(mapIt), inserted}; +} + +template +template +std::pair::iterator, bool> +TCompactMap::TryEmplaceImpl(TKeyParam&& key, TArgs&&... args) +{ + if (IsSmall()) { + auto [arrayIt, found] = ArrayLowerBound(key); + + if (found) { + return {iterator(arrayIt), false}; + } + + if (ArraySize_ < N) { + auto inserted = EmplaceSmall( + arrayIt, + std::piecewise_construct, + std::forward_as_tuple(std::forward(key)), + std::forward_as_tuple(std::forward(args)...)); + return {iterator(inserted), true}; + } + + UpgradeToMap(); + } + + auto [mapIt, inserted] = Map_.try_emplace(std::forward(key), std::forward(args)...); + return {iterator(mapIt), inserted}; +} + +template +auto TCompactMap::erase(const key_type& key) -> size_type +{ + if (IsSmall()) { + auto [arrayIt, found] = ArrayLowerBound(key); + + if (!found) { + return 0; + } + + EraseSmall(arrayIt); + return 1; + } + + return Map_.erase(key); +} + +template +auto TCompactMap::erase(iterator pos) -> iterator +{ + return erase(const_iterator(pos)); +} + +template +auto TCompactMap::erase(const_iterator pos) -> iterator +{ + if (IsSmall()) { + return iterator(EraseSmall(pos.VIter)); + } + + return iterator(Map_.erase(pos.MIter)); +} + +template +auto TCompactMap::erase(const_iterator first, const_iterator last) -> iterator +{ + if (IsSmall()) { + return iterator(EraseSmall(first.VIter, last.VIter)); + } + + return iterator(Map_.erase(first.MIter, last.MIter)); +} + +template +typename TCompactMap::mapped_type& +TCompactMap::operator[](const key_type& key) +{ + return Subscript(key); +} + +template +typename TCompactMap::mapped_type& +TCompactMap::operator[](key_type&& key) +{ + return Subscript(std::move(key)); +} + +template +typename TCompactMap::mapped_type& +TCompactMap::at(const key_type& key) +{ + return AtImpl(this, key); +} + +template +const typename TCompactMap::mapped_type& +TCompactMap::at(const key_type& key) const +{ + return AtImpl(this, key); +} + +//////////////////////////////////////////////////////////////////////////////// + +} // namespace NYT diff --git a/library/cpp/yt/compact_containers/compact_map.h b/library/cpp/yt/compact_containers/compact_map.h new file mode 100644 index 00000000000..4a7443e32c1 --- /dev/null +++ b/library/cpp/yt/compact_containers/compact_map.h @@ -0,0 +1,200 @@ +#pragma once + +#include "comparison.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace NYT { + +//////////////////////////////////////////////////////////////////////////////// + +//! A flat map that keeps up to N elements inline in a sorted buffer, +//! similar to TCompactFlatMap, but, unlike TCompactFlatMap, transparently +//! falls back to std::map storage when the size exceeds N. +//! +//! The fallback is one-way: once the size exceeds N the container keeps using +//! std::map even after it shrinks back; only clear() returns it to inline storage. +//! +//! If switching to std::map storage throws, the container remains valid and +//! retains its keys, but its mapped values may be left in a moved-from state. +template +class TCompactMap +{ +public: + using key_type = TKey; + using mapped_type = TValue; + using value_type = std::pair; + using key_compare = TCompare; + using size_type = std::size_t; + + class iterator; + class const_iterator; + + TCompactMap() = default; + TCompactMap(const TCompactMap& other); + TCompactMap(TCompactMap&& other) noexcept( + std::is_nothrow_move_constructible_v && + std::is_nothrow_move_constructible_v); + TCompactMap& operator=(const TCompactMap& other); + TCompactMap& operator=(TCompactMap&& other) noexcept( + std::is_nothrow_move_constructible_v && + std::is_nothrow_move_assignable_v); + ~TCompactMap(); + + template + TCompactMap(TIt first, TIt last); + + TCompactMap(std::initializer_list init); + + [[nodiscard]] bool empty() const; + size_type size() const; + + void clear(); + + iterator begin(); + const_iterator begin() const; + const_iterator cbegin() const; + + iterator end(); + const_iterator end() const; + const_iterator cend() const; + + template TOtherKey> + size_type count(const TOtherKey& key) const; + + template TOtherKey> + iterator find(const TOtherKey& key); + template TOtherKey> + const_iterator find(const TOtherKey& key) const; + + template TOtherKey> + bool contains(const TOtherKey& key) const; + + std::pair insert(const value_type& value); + std::pair insert(value_type&& value); + + template + std::pair emplace(TArgs&&... args); + + template + std::pair try_emplace(const key_type& key, TArgs&&... args); + + template + std::pair try_emplace(key_type&& key, TArgs&&... args); + + template + void insert(TIt first, TIt last); + + template + std::pair insert_or_assign(const key_type& key, TMapped&& obj); + + size_type erase(const key_type& key); + iterator erase(iterator pos); + iterator erase(const_iterator pos); + iterator erase(const_iterator first, const_iterator last); + + mapped_type& operator[](const key_type& key); + mapped_type& operator[](key_type&& key); + + mapped_type& at(const key_type& key); + const mapped_type& at(const key_type& key) const; + +private: + template + class TIteratorBase; + + using TArrayValue = value_type; + struct alignas(TArrayValue) TStorage + { + std::byte Bytes[sizeof(TArrayValue)]; + }; + using TArrayIterator = TArrayValue*; + using TArrayConstIterator = const TArrayValue*; + using TInlineStorage = std::array; + using TMap = std::map; + using TMapIterator = typename TMap::iterator; + using TMapConstIterator = typename TMap::const_iterator; + + TInlineStorage ArrayStorage_; + size_type ArraySize_ = 0; + TMap Map_; + key_compare Compare_{}; + bool MapMode_ = false; + + TArrayIterator ArrayData(); + TArrayConstIterator ArrayData() const; + TArrayIterator ArrayBegin(); + TArrayConstIterator ArrayBegin() const; + TArrayIterator ArrayEnd(); + TArrayConstIterator ArrayEnd() const; + void ClearArray(); + void CopyArrayFrom(const TCompactMap& other); + void MoveArrayFrom(TCompactMap&& other); + + // Returns true if we are using the inline array. + bool IsSmall() const; + // Moves data from inline array to map. + void UpgradeToMap(); + + // Performs binary search in inline array. Returns iterator and found flag. + template + std::pair ArrayLowerBound(const TOtherKey& key) const; + template + std::pair ArrayLowerBound(const TOtherKey& key); + + template + TArrayIterator EmplaceSmall(TArrayConstIterator pos, TArgs&&... args); + TArrayIterator EraseSmall(TArrayConstIterator pos); + TArrayIterator EraseSmall(TArrayConstIterator first, TArrayConstIterator last); + + template + static std::pair ArrayLowerBoundImpl( + TArrayIteratorType begin, + TArrayIteratorType end, + const TOtherKey& key, + const key_compare& compare); + + template + mapped_type& Subscript(TKeyParam&& key); + + template + static auto BeginImpl(TSelf* self) + -> std::conditional_t, const_iterator, iterator>; + + template + static auto EndImpl(TSelf* self) + -> std::conditional_t, const_iterator, iterator>; + + template + static auto FindImpl(TSelf* self, const TOtherKey& key) + -> std::conditional_t, const_iterator, iterator>; + + template + static auto AtImpl(TSelf* self, const key_type& key) + -> std::conditional_t, const mapped_type&, mapped_type&>; + + template + std::pair TryEmplaceImpl(TKeyParam&& key, TArgs&&... args); + + template + std::pair DoInsert(TValueParam&& value); +}; + +//////////////////////////////////////////////////////////////////////////////// + +} // namespace NYT + +#define COMPACT_MAP_INL_H_ +#include "compact_map-inl.h" +#undef COMPACT_MAP_INL_H_ diff --git a/library/cpp/yt/compact_containers/compact_vector-inl.h b/library/cpp/yt/compact_containers/compact_vector-inl.h index 1070e6b1c91..99798c29d0c 100644 --- a/library/cpp/yt/compact_containers/compact_vector-inl.h +++ b/library/cpp/yt/compact_containers/compact_vector-inl.h @@ -73,8 +73,10 @@ class TCompactVectorReallocationPtrAdjuster template constexpr TCompactVector::TCompactVector() noexcept - : InlineMeta_{} { + if (std::is_constant_evaluated()) { + InlineMeta_ = {}; + } InlineMeta_.SizePlusOne = 1; } @@ -97,7 +99,11 @@ template TCompactVector::TCompactVector(TCompactVector&& other) noexcept(std::is_nothrow_move_constructible_v) : TCompactVector() { - swap(other); + if constexpr (std::is_trivially_copyable_v) { + RelocateFrom(other); + } else { + swap(other); + } } template @@ -435,6 +441,11 @@ void TCompactVector::swap(TCompactVector& other) return; } + if constexpr (std::is_trivially_copyable_v) { + SwapTriviallyCopyable(other); + return; + } + if (!IsInline() && !other.IsInline()) { std::swap(OnHeapMeta_.Storage, other.OnHeapMeta_.Storage); return; @@ -714,6 +725,88 @@ bool TCompactVector::IsInline() const return InlineMeta_.SizePlusOne != 0; } +template +void TCompactVector::RelocateFrom(TCompactVector& other) +{ + static_assert(std::is_trivially_copyable_v); + + if (Y_UNLIKELY(!other.IsInline())) { + OnHeapMeta_.Storage = other.OnHeapMeta_.Storage; + } else if constexpr (PreferFixedSizeMemoryOperations) { + ::memcpy(static_cast(this), static_cast(&other), sizeof(*this)); + } else { + auto sizePlusOne = other.InlineMeta_.SizePlusOne; + ::memcpy(InlineElements_, other.InlineElements_, (sizePlusOne - 1) * sizeof(T)); + InlineMeta_.SizePlusOne = sizePlusOne; + } + + other.InlineMeta_.SizePlusOne = 1; +} + +template +void TCompactVector::SwapTriviallyCopyable(TCompactVector& other) +{ + static_assert(std::is_trivially_copyable_v); + + if constexpr (PreferFixedSizeMemoryOperations) { + if constexpr (sizeof(TCompactVector) > sizeof(TOnHeapStorage*)) { + if (!IsInline() && !other.IsInline()) { + std::swap(OnHeapMeta_.Storage, other.OnHeapMeta_.Storage); + return; + } + } + + static_assert(sizeof(TCompactVector) % sizeof(uintptr_t) == 0); + auto* lhs = reinterpret_cast(this); + auto* rhs = reinterpret_cast(&other); + for (size_t index = 0; index < sizeof(TCompactVector); index += sizeof(uintptr_t)) { + uintptr_t buffer; + ::memcpy(&buffer, lhs + index, sizeof(uintptr_t)); + ::memcpy(lhs + index, rhs + index, sizeof(uintptr_t)); + ::memcpy(rhs + index, &buffer, sizeof(uintptr_t)); + } + return; + } + + bool thisIsInline = IsInline(); + bool otherIsInline = other.IsInline(); + if (!thisIsInline && !otherIsInline) { + std::swap(OnHeapMeta_.Storage, other.OnHeapMeta_.Storage); + return; + } + + if (thisIsInline != otherIsInline) { + auto* inlineVector = thisIsInline ? this : &other; + auto* onHeapVector = thisIsInline ? &other : this; + auto* storage = onHeapVector->OnHeapMeta_.Storage; + auto sizePlusOne = inlineVector->InlineMeta_.SizePlusOne; + ::memcpy( + onHeapVector->InlineElements_, + inlineVector->InlineElements_, + (sizePlusOne - 1) * sizeof(T)); + onHeapVector->InlineMeta_.SizePlusOne = sizePlusOne; + inlineVector->OnHeapMeta_.Storage = storage; + return; + } + + size_t thisSize = InlineMeta_.SizePlusOne - 1; + size_t otherSize = other.InlineMeta_.SizePlusOne - 1; + size_t commonSize = std::min(thisSize, otherSize); + std::swap_ranges(InlineElements_, InlineElements_ + commonSize, other.InlineElements_); + if (thisSize < otherSize) { + ::memcpy( + InlineElements_ + commonSize, + other.InlineElements_ + commonSize, + (otherSize - commonSize) * sizeof(T)); + } else if (otherSize < thisSize) { + ::memcpy( + other.InlineElements_ + commonSize, + InlineElements_ + commonSize, + (thisSize - commonSize) * sizeof(T)); + } + std::swap(InlineMeta_.SizePlusOne, other.InlineMeta_.SizePlusOne); +} + template void TCompactVector::SetSize(size_t newSize) { diff --git a/library/cpp/yt/compact_containers/compact_vector.h b/library/cpp/yt/compact_containers/compact_vector.h index 77311c89f9e..5ecad2ff819 100644 --- a/library/cpp/yt/compact_containers/compact_vector.h +++ b/library/cpp/yt/compact_containers/compact_vector.h @@ -5,6 +5,7 @@ #include #include #include +#include namespace NYT { @@ -152,6 +153,9 @@ class TCompactVector using TOnHeapStorage = TCompactVectorOnHeapStorage; + static constexpr bool PreferFixedSizeMemoryOperations = + sizeof(TCompactVector) <= 8 * sizeof(uintptr_t); + static constexpr size_t ByteSize = (sizeof(T) * N + alignof(T) + sizeof(uintptr_t) - 1) & ~(sizeof(uintptr_t) - 1); @@ -189,6 +193,8 @@ class TCompactVector }; bool IsInline() const; + void RelocateFrom(TCompactVector& other); + void SwapTriviallyCopyable(TCompactVector& other); void SetSize(size_t newSize); void EnsureOnHeapCapacity(size_t newCapacity, bool incremental); template diff --git a/library/cpp/yt/compact_containers/comparison.h b/library/cpp/yt/compact_containers/comparison.h new file mode 100644 index 00000000000..03c98c34775 --- /dev/null +++ b/library/cpp/yt/compact_containers/comparison.h @@ -0,0 +1,19 @@ +#pragma once + +#include + +namespace NYT::NDetail { + +//////////////////////////////////////////////////////////////////////////////// + +template +concept CHasIsTransparentFlag = requires { + typename T::is_transparent; +}; + +template +concept CComparisonAllowed = std::same_as || CHasIsTransparentFlag; + +//////////////////////////////////////////////////////////////////////////////// + +} // namespace NYT::NDetail diff --git a/library/cpp/yt/compact_containers/unittests/compact_map_ut.cpp b/library/cpp/yt/compact_containers/unittests/compact_map_ut.cpp new file mode 100644 index 00000000000..266e6bd13cf --- /dev/null +++ b/library/cpp/yt/compact_containers/unittests/compact_map_ut.cpp @@ -0,0 +1,1386 @@ +#include + +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace NYT { +namespace { + +//////////////////////////////////////////////////////////////////////////////// + +using TMap = TCompactMap; + +TMap CreateMap() +{ + std::vector> data = { + {"I", "met"}, + {"a", "traveller"}, + {"from", "an"}, + {"antique", "land"}, + }; + return {data.begin(), data.end()}; +} + +//////////////////////////////////////////////////////////////////////////////// +// Basic operations. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, DefaultEmpty) +{ + TMap m; + EXPECT_TRUE(m.empty()); + EXPECT_EQ(m.begin(), m.end()); + EXPECT_EQ(m.size(), 0u); +} + +TEST(TCompactMapTest, Size) +{ + auto m = CreateMap(); + + EXPECT_EQ(m.size(), 4u); + + m.insert({"Who", "said"}); + + EXPECT_EQ(m.size(), 5u); + + m.erase("antique"); + + EXPECT_EQ(m.size(), 4u); +} + +TEST(TCompactMapTest, ClearAndEmpty) +{ + auto m = CreateMap(); + + EXPECT_FALSE(m.empty()); + EXPECT_NE(m.begin(), m.end()); + + m.clear(); + + EXPECT_TRUE(m.empty()); + EXPECT_EQ(m.begin(), m.end()); + + m.insert({"Who", "said"}); + + EXPECT_FALSE(m.empty()); + EXPECT_NE(m.begin(), m.end()); +} + +//////////////////////////////////////////////////////////////////////////////// +// Small size (inline vector) tests. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, InsertSmall) +{ + TCompactMap m; + + auto [it1, inserted1] = m.insert({1, "one"}); + EXPECT_TRUE(inserted1); + EXPECT_EQ(it1->first, 1); + EXPECT_EQ(it1->second, "one"); + EXPECT_EQ(m.size(), 1u); + + auto [it2, inserted2] = m.insert({2, "two"}); + EXPECT_TRUE(inserted2); + EXPECT_EQ(m.size(), 2u); + + auto [it3, inserted3] = m.insert({1, "uno"}); + EXPECT_FALSE(inserted3); + // The value is not changed. + EXPECT_EQ(it3->second, "one"); + EXPECT_EQ(m.size(), 2u); +} + +TEST(TCompactMapTest, InsertMoveOnlyValue) +{ + TCompactMap, 1> m; + + auto [smallIt, smallInserted] = m.insert({1, std::make_unique(10)}); + EXPECT_TRUE(smallInserted); + EXPECT_EQ(*smallIt->second, 10); + + auto [largeIt, largeInserted] = m.insert({2, std::make_unique(20)}); + EXPECT_TRUE(largeInserted); + EXPECT_EQ(*largeIt->second, 20); +} + +TEST(TCompactMapTest, FindSmall) +{ + TCompactMap m; + m.insert({1, "one"}); + m.insert({2, "two"}); + m.insert({3, "three"}); + + auto it = m.find(2); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it->first, 2); + EXPECT_EQ(it->second, "two"); + + auto it2 = m.find(99); + EXPECT_EQ(it2, m.end()); +} + +TEST(TCompactMapTest, EraseSmall) +{ + TCompactMap m; + m.insert({1, "one"}); + m.insert({2, "two"}); + m.insert({3, "three"}); + + EXPECT_TRUE(m.erase(2)); + EXPECT_EQ(m.size(), 2u); + EXPECT_EQ(m.find(2), m.end()); + EXPECT_NE(m.find(1), m.end()); + EXPECT_NE(m.find(3), m.end()); + + EXPECT_FALSE(m.erase(99)); + EXPECT_EQ(m.size(), 2u); +} + +TEST(TCompactMapTest, IteratorSmall) +{ + TCompactMap m; + m.insert({3, "three"}); + m.insert({1, "one"}); + m.insert({2, "two"}); + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + std::sort(keys.begin(), keys.end()); + EXPECT_EQ(keys, std::vector({1, 2, 3})); +} + +//////////////////////////////////////////////////////////////////////////////// +// Large size (std::map) tests. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, GrowToLarge) +{ + TCompactMap m; + + for (int i = 0; i < 8; i++) { + m.insert({i, std::to_string(i)}); + } + + EXPECT_EQ(m.size(), 8u); + + for (int i = 0; i < 8; i++) { + EXPECT_EQ(m.count(i), 1u); + EXPECT_EQ(m.find(i)->second, std::to_string(i)); + } + + EXPECT_EQ(m.count(99), 0u); +} + +TEST(TCompactMapTest, InsertLarge) +{ + TCompactMap m; + + for (int i = 0; i < 5; i++) { + auto [it, inserted] = m.insert({i, std::to_string(i)}); + EXPECT_TRUE(inserted); + EXPECT_EQ(it->first, i); + EXPECT_EQ(it->second, std::to_string(i)); + } + + EXPECT_EQ(m.size(), 5u); + + auto [it, inserted] = m.insert({2, "two"}); + EXPECT_FALSE(inserted); + // The value is not changed. + EXPECT_EQ(it->second, "2"); + EXPECT_EQ(m.size(), 5u); +} + +TEST(TCompactMapTest, FindLarge) +{ + TCompactMap m; + + for (int i = 0; i < 8; i++) { + m.insert({i, std::to_string(i)}); + } + + for (int i = 0; i < 8; i++) { + auto it = m.find(i); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it->first, i); + EXPECT_EQ(it->second, std::to_string(i)); + } + + EXPECT_EQ(m.find(99), m.end()); +} + +TEST(TCompactMapTest, EraseLarge) +{ + TCompactMap m; + + for (int i = 0; i < 8; i++) { + m.insert({i, std::to_string(i)}); + } + + for (int i = 0; i < 8; i++) { + EXPECT_EQ(m.count(i), 1u); + EXPECT_TRUE(m.erase(i)); + EXPECT_EQ(m.count(i), 0u); + EXPECT_EQ(m.size(), 8u - i - 1); + + for (int j = i + 1; j < 8; j++) { + EXPECT_EQ(m.count(j), 1u); + } + } + + EXPECT_EQ(m.count(99), 0u); +} + +TEST(TCompactMapTest, IteratorLarge) +{ + TCompactMap m; + + for (int i = 0; i < 6; i++) { + m.insert({i, std::to_string(i)}); + } + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + std::sort(keys.begin(), keys.end()); + std::vector expected = {0, 1, 2, 3, 4, 5}; + EXPECT_EQ(keys, expected); +} + +//////////////////////////////////////////////////////////////////////////////// +// Transition from small to large. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, SmallToLargeTransition) +{ + TCompactMap m; + + // Fill small storage + for (int i = 0; i < 4; i++) { + m.insert({i, std::to_string(i)}); + } + + EXPECT_EQ(m.size(), 4u); + + // Add one more to trigger transition. + m.insert({4, "4"}); + + EXPECT_EQ(m.size(), 5u); + + // Check all elements are still there. + for (int i = 0; i < 5; i++) { + EXPECT_EQ(m.count(i), 1u); + EXPECT_EQ(m.find(i)->second, std::to_string(i)); + } +} + +TEST(TCompactMapTest, GrowShrink) +{ + TCompactMap m; + m.insert({"Two", "vast"}); + m.insert({"and", "trunkless"}); + m.insert({"legs", "of"}); + m.insert({"stone", "Stand"}); + m.insert({"in", "the"}); + m.insert({"desert", "..."}); + + EXPECT_EQ(m.size(), 6u); + + m.erase("legs"); + m.erase("stone"); + m.erase("in"); + m.erase("desert"); + + EXPECT_EQ(m.size(), 2u); + + // All remaining elements should be accessible. + EXPECT_NE(m.find("Two"), m.end()); + EXPECT_NE(m.find("and"), m.end()); +} + +//////////////////////////////////////////////////////////////////////////////// +// Find operations (mutable and const). +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, FindMutable) +{ + auto m = CreateMap(); + { + auto it = m.find("from"); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it->second, "an"); + it->second = "the"; + } + { + auto it = m.find("from"); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it->second, "the"); + } + { + auto it = m.find("Who"); + EXPECT_EQ(it, m.end()); + } +} + +TEST(TCompactMapTest, FindConst) +{ + const auto& m = CreateMap(); + { + auto it = m.find("from"); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it->second, "an"); + } + { + auto it = m.find("Who"); + EXPECT_EQ(it, m.end()); + } +} + +TEST(TCompactMapTest, Contains) +{ + auto m = CreateMap(); + + EXPECT_TRUE(m.contains("I")); + EXPECT_TRUE(m.contains("from")); + EXPECT_FALSE(m.contains("Who")); + + m.erase("I"); + EXPECT_FALSE(m.contains("I")); +} + +TEST(TCompactMapTest, HeterogeneousLookupSmallAndLarge) +{ + using namespace std::string_view_literals; + + TCompactMap small = { + {"alpha", 1}, + {"beta", 2}, + }; + + EXPECT_NE(small.find("alpha"sv), small.end()); + EXPECT_TRUE(small.contains("beta"sv)); + EXPECT_EQ(small.count("missing"sv), 0u); + + TCompactMap large = { + {"alpha", 1}, + {"beta", 2}, + {"gamma", 3}, + }; + + EXPECT_NE(large.find("alpha"sv), large.end()); + EXPECT_TRUE(large.contains("gamma"sv)); + EXPECT_EQ(large.count("missing"sv), 0u); +} + +//////////////////////////////////////////////////////////////////////////////// +// Insert operations. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, Insert) +{ + auto m = CreateMap(); + + auto [it, inserted] = m.insert({"Who", "said"}); + EXPECT_TRUE(inserted); + EXPECT_EQ(m.size(), 5u); + EXPECT_NE(it, m.end()); + EXPECT_EQ(it, m.find("Who")); + EXPECT_EQ(it->second, "said"); + + auto [it2, inserted2] = m.insert({"Who", "told"}); + EXPECT_FALSE(inserted2); + EXPECT_EQ(m.size(), 5u); + EXPECT_EQ(it2, it); + EXPECT_EQ(it->second, "said"); + + std::vector> data = { + {"Two", "vast"}, + {"and", "trunkless"}, + {"legs", "of"}, + }; + m.insert(data.begin(), data.end()); + EXPECT_EQ(m.size(), 8u); + EXPECT_NE(m.find("and"), m.end()); + EXPECT_EQ(m.find("and")->second, "trunkless"); +} + +TEST(TCompactMapTest, Emplace) +{ + TCompactMap m; + + auto [it1, inserted1] = m.emplace(1, "one"); + EXPECT_TRUE(inserted1); + EXPECT_EQ(it1->second, "one"); + + auto [it2, inserted2] = m.emplace(1, "uno"); + EXPECT_FALSE(inserted2); + EXPECT_EQ(it2->second, "one"); + + auto [it3, inserted3] = m.emplace(0, "zero"); + EXPECT_TRUE(inserted3); + EXPECT_EQ(it3->first, 0); + EXPECT_EQ(m.size(), 2u); + + for (int i = 2; i < 6; ++i) { + m.emplace(i, std::to_string(i)); + } + + EXPECT_EQ(m.size(), 6u); + auto [tailIt, tailInserted] = m.emplace(10, "ten"); + EXPECT_TRUE(tailInserted); + EXPECT_EQ(tailIt->first, 10); + + auto [dupIt, dupInserted] = m.emplace(3, "tres"); + EXPECT_FALSE(dupInserted); + EXPECT_EQ(dupIt->second, "3"); +} + +TEST(TCompactMapTest, TryEmplace) +{ + TCompactMap m; + + auto [it1, inserted1] = m.try_emplace(1, "one"); + EXPECT_TRUE(inserted1); + EXPECT_EQ(it1->second, "one"); + + auto [it2, inserted2] = m.try_emplace(1, "uno"); + EXPECT_FALSE(inserted2); + EXPECT_EQ(it2->second, "one"); + + auto [it3, inserted3] = m.try_emplace(0, "zero"); + EXPECT_TRUE(inserted3); + EXPECT_EQ(it3->first, 0); + + auto [it4, inserted4] = m.try_emplace(5, "five"); + EXPECT_TRUE(inserted4); + + // Grow to std::map backed mode. + for (int i = 6; i < 9; ++i) { + m.try_emplace(i, std::to_string(i)); + } + + auto [it5, inserted5] = m.try_emplace(9, "nine"); + EXPECT_TRUE(inserted5); + + auto [it6, inserted6] = m.try_emplace(5, "FIVE"); + EXPECT_FALSE(inserted6); + EXPECT_EQ(it6->second, "five"); +} + +TEST(TCompactMapTest, InsertOrAssign) +{ + TCompactMap m; + + auto [it1, inserted1] = m.insert_or_assign(1, "one"); + EXPECT_TRUE(inserted1); + EXPECT_EQ(it1->second, "one"); + EXPECT_EQ(m.size(), 1u); + + auto [it2, inserted2] = m.insert_or_assign(1, "uno"); + EXPECT_FALSE(inserted2); + EXPECT_EQ(it2->second, "uno"); + EXPECT_EQ(m.size(), 1u); + + // Test after growing to large. + for (int i = 2; i < 6; i++) { + m.insert_or_assign(i, std::to_string(i)); + } + + auto [it3, inserted3] = m.insert_or_assign(3, "three"); + EXPECT_FALSE(inserted3); + EXPECT_EQ(it3->second, "three"); +} + +TEST(TCompactMapTest, InitializerList) +{ + TCompactMap m = { + {1, "one"}, + {2, "two"}, + {3, "three"}, + }; + + EXPECT_EQ(m.size(), 3u); + EXPECT_EQ(m.find(1)->second, "one"); + EXPECT_EQ(m.find(2)->second, "two"); + EXPECT_EQ(m.find(3)->second, "three"); +} + +//////////////////////////////////////////////////////////////////////////////// +// Subscript operator +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, SubscriptSmall) +{ + TCompactMap m; + + m["key1"] = "value1"; + EXPECT_EQ(m["key1"], "value1"); + EXPECT_EQ(m.size(), 1u); + + m["key2"] = "value2"; + EXPECT_EQ(m.size(), 2u); + + EXPECT_EQ(m["key3"], ""); + EXPECT_EQ(m.size(), 3u); +} + +TEST(TCompactMapTest, SubscriptLarge) +{ + TCompactMap m; + + for (int i = 0; i < 5; i++) { + m["key" + std::to_string(i)] = "value" + std::to_string(i); + } + + EXPECT_EQ(m.size(), 5u); + EXPECT_EQ(m["key3"], "value3"); + + m["key3"] = "new_value"; + EXPECT_EQ(m["key3"], "new_value"); +} + +TEST(TCompactMapTest, SubscriptRvalue) +{ + TCompactMap m; + + std::string key = "temp"; + m[std::move(key)] = "value"; + EXPECT_EQ(m["temp"], "value"); +} + +//////////////////////////////////////////////////////////////////////////////// +// At operations. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, At) +{ + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + + EXPECT_EQ(m.at(1), "one"); + EXPECT_EQ(m.at(2), "two"); + + EXPECT_THROW(m.at(99), std::out_of_range); + + m.at(1) = "uno"; + EXPECT_EQ(m.at(1), "uno"); +} + +TEST(TCompactMapTest, AtConst) +{ + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + + const auto& cm = m; + + EXPECT_EQ(cm.at(1), "one"); + EXPECT_EQ(cm.at(2), "two"); + + EXPECT_THROW(cm.at(99), std::out_of_range); +} + +TEST(TCompactMapTest, AtLarge) +{ + TCompactMap m; + for (int i = 0; i < 6; i++) { + m[i] = std::to_string(i); + } + + EXPECT_EQ(m.at(3), "3"); + EXPECT_THROW(m.at(99), std::out_of_range); +} + +//////////////////////////////////////////////////////////////////////////////// +// Iterator tests. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, IteratorModification) +{ + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + + for (auto it = m.begin(); it != m.end(); ++it) { + it->second += "_modified"; + } + + EXPECT_EQ(m[1], "one_modified"); + EXPECT_EQ(m[2], "two_modified"); + EXPECT_EQ(m[3], "three_modified"); +} + +TEST(TCompactMapTest, IteratorString) +{ + TCompactMap m; + + m["str1"] = "val1"; + m["str2"] = "val2"; + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + std::sort(keys.begin(), keys.end()); + EXPECT_EQ(m.size(), 2u); + EXPECT_EQ(keys[0], "str1"); + EXPECT_EQ(keys[1], "str2"); + + // Add more to grow. + m["str4"] = "val4"; + m["str0"] = "val0"; + + keys.clear(); + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + std::sort(keys.begin(), keys.end()); + EXPECT_EQ(m.size(), 4u); + EXPECT_EQ(keys[0], "str0"); + EXPECT_EQ(keys[1], "str1"); + EXPECT_EQ(keys[2], "str2"); + EXPECT_EQ(keys[3], "str4"); +} + +TEST(TCompactMapTest, IteratorConversion) +{ + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + + auto it = m.begin(); + TCompactMap::const_iterator cit = it; + + EXPECT_EQ(*it, *cit); + EXPECT_EQ(it, cit); + EXPECT_EQ(cit, it); + EXPECT_NE(it, m.cend()); + EXPECT_NE(m.cend(), it); +} + +TEST(TCompactMapTest, IteratorTypeProperties) +{ + using TMap = TCompactMap; + using TIter = TMap::iterator; + using TCIter = TMap::const_iterator; + + static_assert(std::is_default_constructible_v); + static_assert(std::is_default_constructible_v); + static_assert(std::bidirectional_iterator); + static_assert(std::bidirectional_iterator); + + // Key type cannot be modified via any iterator. + static_assert(!std::is_assignable_v()->first)), std::string>); + static_assert(!std::is_assignable_v()->first)), std::string>); + + // Mapped type can only be modified via non-const iterator. + static_assert(std::is_assignable_v()->second)), std::string>); + static_assert(!std::is_assignable_v()->second)), std::string>); +} + +TEST(TCompactMapTest, ErasePostIncrementedIterator) +{ + TCompactMap map = { + {1, 10}, + {2, 20}, + }; + + auto it = map.begin(); + auto next = map.erase(it++); + + ASSERT_EQ(map.size(), 1u); + EXPECT_EQ(next, map.begin()); + EXPECT_EQ(next->first, 2); +} + +TEST(TCompactMapTest, IteratorDecrementSmallAndLarge) +{ + { + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + + auto it = m.end(); + --it; + EXPECT_EQ(it->first, 3); + EXPECT_EQ((it--)->first, 3); + EXPECT_EQ(it->first, 2); + --it; + EXPECT_EQ(it->first, 1); + EXPECT_EQ(it, m.begin()); + } + + { + TCompactMap m; + for (int i = 0; i < 5; ++i) { + m[i] = std::to_string(i); + } + + auto it = m.end(); + --it; + EXPECT_EQ(it->first, 4); + EXPECT_EQ((it--)->first, 4); + EXPECT_EQ(it->first, 3); + --it; + EXPECT_EQ(it->first, 2); + --it; + EXPECT_EQ(it->first, 1); + --it; + EXPECT_EQ(it->first, 0); + EXPECT_EQ(it, m.begin()); + } +} + +TEST(TCompactMapTest, EraseByIteratorSmallAndLarge) +{ + { + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + + auto middle = m.find(2); + ASSERT_NE(middle, m.end()); + auto next = m.erase(middle); + EXPECT_EQ(next->first, 3); + EXPECT_EQ(m.count(2), 0u); + + auto rangeBegin = m.begin(); + auto rangeEnd = m.end(); + // Removing the tail via iterator range should leave the container empty. + auto afterRange = m.erase(rangeBegin, rangeEnd); + EXPECT_EQ(afterRange, m.end()); + EXPECT_TRUE(m.empty()); + } + + { + TCompactMap m; + for (int i = 0; i < 5; ++i) { + m[i] = std::to_string(i); + } + + auto it = m.find(2); + ASSERT_NE(it, m.end()); + auto next = m.erase(it); + EXPECT_EQ(next->first, 3); + EXPECT_EQ(m.count(2), 0u); + + auto first = m.find(0); + auto last = m.find(3); + ASSERT_NE(first, m.end()); + ASSERT_NE(last, m.end()); + auto afterRange = m.erase(first, last); + EXPECT_EQ(afterRange->first, 3); + std::vector remaining; + for (const auto& [k, v] : m) { + remaining.push_back(k); + } + EXPECT_EQ(remaining, std::vector({3, 4})); + } +} + +TEST(TCompactMapTest, EraseUpgradedMapToEndOneByOne) +{ + TCompactMap map = { + {1, 1}, + {2, 2}, + {3, 3}, + }; + + for (auto it = map.begin(); it != map.end();) { + it = map.erase(it); + } + + EXPECT_TRUE(map.empty()); + EXPECT_EQ(map.begin(), map.end()); +} + +TEST(TCompactMapTest, EraseEntireUpgradedMapByRange) +{ + TCompactMap map = { + {1, 1}, + {2, 2}, + {3, 3}, + }; + + auto result = map.erase(map.begin(), map.end()); + + EXPECT_EQ(result, map.end()); + EXPECT_TRUE(map.empty()); +} + +TEST(TCompactMapTest, StructuredBindingModifySmallAndLarge) +{ + // Small-mode: everything stays in the inline storage. + { + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + + for (auto&& [k, v] : m) { + v += "_x"; + } + + EXPECT_EQ(m[1], "one_x"); + EXPECT_EQ(m[2], "two_x"); + EXPECT_EQ(m[3], "three_x"); + } + + // Large-mode: container has switched to std::map storage. + { + TCompactMap m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + // Triggers growth to large mode. + m[4] = "four"; + + for (auto&& [k, v] : m) { + v += "_y"; + } + + EXPECT_EQ(m[1], "one_y"); + EXPECT_EQ(m[2], "two_y"); + EXPECT_EQ(m[3], "three_y"); + EXPECT_EQ(m[4], "four_y"); + } +} + +//////////////////////////////////////////////////////////////////////////////// +// Edge cases +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, EmptyAfterClear) +{ + TCompactMap m; + + for (int i = 0; i < 8; i++) { + m[i] = std::to_string(i); + } + + m.clear(); + + EXPECT_TRUE(m.empty()); + EXPECT_EQ(m.size(), 0u); + EXPECT_EQ(m.begin(), m.end()); +} + +TEST(TCompactMapTest, CountZeroOrOne) +{ + TCompactMap m; + + EXPECT_EQ(m.count(1), 0u); + + m[1] = "one"; + EXPECT_EQ(m.count(1), 1u); + + m.erase(1); + EXPECT_EQ(m.count(1), 0u); +} + +TEST(TCompactMapTest, SortedOrder) +{ + TCompactMap m; + m[5] = "five"; + m[1] = "one"; + m[3] = "three"; + m[2] = "two"; + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + EXPECT_EQ(keys, std::vector({1, 2, 3, 5})); +} + +TEST(TCompactMapTest, CustomComparator) +{ + TCompactMap> m; + m[1] = "one"; + m[2] = "two"; + m[3] = "three"; + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + EXPECT_EQ(keys, std::vector({3, 2, 1})); +} + +TEST(TCompactMapTest, CustomComparatorAfterUpgrade) +{ + TCompactMap> m; + for (int i = 0; i < 5; ++i) { + m[i] = std::to_string(i); + } + + std::vector keys; + for (const auto& [k, v] : m) { + keys.push_back(k); + } + + EXPECT_EQ(keys, std::vector({4, 3, 2, 1, 0})); +} + +TEST(TCompactMapTest, ZeroInlineCapacityBehavesLikeMap) +{ + TCompactMap m; + EXPECT_TRUE(m.empty()); + + m[7] = "seven"; + EXPECT_EQ(m.size(), 1u); + EXPECT_EQ(m.at(7), "seven"); + + auto [it, inserted] = m.insert_or_assign(7, "SEVEN"); + EXPECT_FALSE(inserted); + EXPECT_EQ(it->second, "SEVEN"); + + auto [it2, inserted2] = m.insert_or_assign(4, "four"); + EXPECT_TRUE(inserted2); + EXPECT_EQ(m.size(), 2u); + EXPECT_EQ(m.at(4), "four"); + + EXPECT_EQ(m.erase(7), 1u); + EXPECT_EQ(m.erase(7), 0u); + EXPECT_FALSE(m.contains(7)); + EXPECT_EQ(m.begin()->first, 4); +} + +//////////////////////////////////////////////////////////////////////////////// +// Exception safety. +//////////////////////////////////////////////////////////////////////////////// + +struct TThrowingCopy +{ + static int Alive; + static int CopiesBeforeThrow; + + int Value; + + explicit TThrowingCopy(int value) + : Value(value) + { + ++Alive; + } + + TThrowingCopy(const TThrowingCopy& other) + : Value(other.Value) + { + if (CopiesBeforeThrow == 0) { + throw std::runtime_error("copy failed"); + } + if (CopiesBeforeThrow > 0) { + --CopiesBeforeThrow; + } + ++Alive; + } + + TThrowingCopy(TThrowingCopy&& other) noexcept + : Value(other.Value) + { + ++Alive; + } + + TThrowingCopy& operator=(const TThrowingCopy&) = default; + TThrowingCopy& operator=(TThrowingCopy&&) = default; + + ~TThrowingCopy() + { + --Alive; + } +}; + +int TThrowingCopy::Alive = 0; +int TThrowingCopy::CopiesBeforeThrow = -1; + +TEST(TCompactMapTest, FailedInlineCopyDestroysPartialCopy) +{ + using TTestMap = TCompactMap; + + EXPECT_EQ(TThrowingCopy::Alive, 0); + { + TTestMap source; + source.try_emplace(1, 10); + source.try_emplace(2, 20); + source.try_emplace(3, 30); + + TThrowingCopy::CopiesBeforeThrow = 1; + EXPECT_THROW(TTestMap copy(source), std::runtime_error); + EXPECT_EQ(TThrowingCopy::Alive, 3); + + TTestMap target; + target.try_emplace(4, 40); + TThrowingCopy::CopiesBeforeThrow = 1; + EXPECT_THROW(target = source, std::runtime_error); + EXPECT_TRUE(target.empty()); + EXPECT_EQ(TThrowingCopy::Alive, 3); + } + EXPECT_EQ(TThrowingCopy::Alive, 0); + + TThrowingCopy::CopiesBeforeThrow = -1; +} + +struct TThrowingMapped +{ + static int Alive; + static int MovesBeforeThrow; + static bool ThrowOnConstruction; + + int Value; + + explicit TThrowingMapped(int value) + : Value(value) + { + if (ThrowOnConstruction) { + throw std::runtime_error("construction failed"); + } + ++Alive; + } + + TThrowingMapped(const TThrowingMapped&) = delete; + TThrowingMapped& operator=(const TThrowingMapped&) = delete; + + TThrowingMapped(TThrowingMapped&& other) // NOLINT + : Value(other.Value) + { + if (MovesBeforeThrow == 0) { + throw std::runtime_error("move failed"); + } + if (MovesBeforeThrow > 0) { + --MovesBeforeThrow; + } + ++Alive; + } + + TThrowingMapped& operator=(TThrowingMapped&&) = delete; + + ~TThrowingMapped() + { + --Alive; + } +}; + +int TThrowingMapped::Alive = 0; +int TThrowingMapped::MovesBeforeThrow = -1; +bool TThrowingMapped::ThrowOnConstruction = false; + +TEST(TCompactMapTest, FailedInlineValueConstructionLeavesMapUnchanged) +{ + TCompactMap map; + map.try_emplace(1, 10); + map.try_emplace(3, 30); + + TThrowingMapped::ThrowOnConstruction = true; + EXPECT_THROW(map.try_emplace(2, 20), std::runtime_error); + TThrowingMapped::ThrowOnConstruction = false; + + EXPECT_EQ(map.size(), 2u); + EXPECT_EQ(map.at(1).Value, 10); + EXPECT_EQ(map.at(3).Value, 30); +} + +TEST(TCompactMapTest, FailedInlineRelocationLeavesValidMap) +{ + EXPECT_EQ(TThrowingMapped::Alive, 0); + { + TCompactMap map; + map.try_emplace(1, 10); + map.try_emplace(2, 20); + map.try_emplace(3, 30); + + TThrowingMapped::MovesBeforeThrow = 1; + EXPECT_THROW(map.try_emplace(0, 0), std::runtime_error); + TThrowingMapped::MovesBeforeThrow = -1; + + EXPECT_EQ(map.size(), 2u); + EXPECT_EQ(map.at(1).Value, 10); + EXPECT_EQ(map.at(2).Value, 20); + + map.try_emplace(4, 40); + EXPECT_EQ(map.at(4).Value, 40); + } + EXPECT_EQ(TThrowingMapped::Alive, 0); +} + +TEST(TCompactMapTest, FailedInlineMoveDestroysPartialMove) +{ + using TTestMap = TCompactMap; + + EXPECT_EQ(TThrowingMapped::Alive, 0); + { + TTestMap source; + source.try_emplace(1, 10); + source.try_emplace(2, 20); + source.try_emplace(3, 30); + + TThrowingMapped::MovesBeforeThrow = 1; + EXPECT_THROW(TTestMap moved(std::move(source)), std::runtime_error); + EXPECT_EQ(TThrowingMapped::Alive, 3); + + TTestMap target; + target.try_emplace(4, 40); + TThrowingMapped::MovesBeforeThrow = 1; + EXPECT_THROW(target = std::move(source), std::runtime_error); + EXPECT_TRUE(target.empty()); + EXPECT_EQ(TThrowingMapped::Alive, 3); + } + EXPECT_EQ(TThrowingMapped::Alive, 0); + + TThrowingMapped::MovesBeforeThrow = -1; +} + +TEST(TCompactMapTest, FailedInlineEraseRelocationLeavesValidMap) +{ + EXPECT_EQ(TThrowingMapped::Alive, 0); + { + TCompactMap map; + map.try_emplace(1, 10); + map.try_emplace(2, 20); + map.try_emplace(3, 30); + map.try_emplace(4, 40); + + TThrowingMapped::MovesBeforeThrow = 1; + EXPECT_THROW(map.erase(map.begin()), std::runtime_error); + TThrowingMapped::MovesBeforeThrow = -1; + + ASSERT_EQ(map.size(), 1u); + EXPECT_EQ(map.begin()->first, 2); + EXPECT_EQ(map.begin()->second.Value, 20); + + map.try_emplace(5, 50); + EXPECT_EQ(map.at(5).Value, 50); + } + EXPECT_EQ(TThrowingMapped::Alive, 0); +} + +//////////////////////////////////////////////////////////////////////////////// +// Copy, move, and lifetime management. +//////////////////////////////////////////////////////////////////////////////// + +TEST(TCompactMapTest, CopyAndMoveConstructionSmallAndLarge) +{ + using TTestMap = TCompactMap; + + TTestMap small = { + {1, 10}, + {2, 20}, + }; + TTestMap smallCopy(small); + EXPECT_EQ(smallCopy.size(), 2u); + EXPECT_EQ(smallCopy.at(1), 10); + EXPECT_EQ(smallCopy.at(2), 20); + + TTestMap smallMove(std::move(smallCopy)); + EXPECT_EQ(smallMove.size(), 2u); + EXPECT_TRUE(smallCopy.empty()); + smallCopy.emplace(3, 30); + EXPECT_EQ(smallCopy.at(3), 30); + + TTestMap large = { + {1, 10}, + {2, 20}, + {3, 30}, + }; + TTestMap largeCopy(large); + EXPECT_EQ(largeCopy.size(), 3u); + EXPECT_EQ(largeCopy.at(3), 30); + + TTestMap largeMove(std::move(largeCopy)); + EXPECT_EQ(largeMove.size(), 3u); + EXPECT_TRUE(largeCopy.empty()); + largeCopy.emplace(4, 40); + EXPECT_EQ(largeCopy.at(4), 40); +} + +TEST(TCompactMapTest, CopyAndMoveAssignmentBetweenSmallAndLarge) +{ + using TTestMap = TCompactMap; + + TTestMap smallSource = {{1, 10}}; + TTestMap largeTarget = { + {2, 20}, + {3, 30}, + {4, 40}, + }; + largeTarget = smallSource; + EXPECT_EQ(largeTarget.size(), 1u); + EXPECT_EQ(largeTarget.at(1), 10); + + TTestMap largeSource = { + {5, 50}, + {6, 60}, + {7, 70}, + }; + TTestMap smallTarget = {{8, 80}}; + smallTarget = largeSource; + EXPECT_EQ(smallTarget.size(), 3u); + EXPECT_EQ(smallTarget.at(7), 70); + + TTestMap smallMoveSource = {{9, 90}}; + largeTarget = std::move(smallMoveSource); + EXPECT_EQ(largeTarget.at(9), 90); + EXPECT_TRUE(smallMoveSource.empty()); + smallMoveSource.emplace(10, 100); + EXPECT_EQ(smallMoveSource.at(10), 100); + + TTestMap largeMoveSource = { + {11, 110}, + {12, 120}, + {13, 130}, + }; + smallTarget = std::move(largeMoveSource); + EXPECT_EQ(smallTarget.size(), 3u); + EXPECT_EQ(smallTarget.at(13), 130); + EXPECT_TRUE(largeMoveSource.empty()); + largeMoveSource.emplace(14, 140); + EXPECT_EQ(largeMoveSource.at(14), 140); +} + +TEST(TCompactMapTest, SelfAssignmentSmallAndLarge) +{ + auto copyAssign = [] (auto* lhs, const auto* rhs) { + *lhs = *rhs; + }; + auto moveAssign = [] (auto* lhs, auto* rhs) { + *lhs = std::move(*rhs); + }; + + TCompactMap small = {{1, 10}}; + copyAssign(&small, &small); + EXPECT_EQ(small.at(1), 10); + moveAssign(&small, &small); + EXPECT_EQ(small.at(1), 10); + + TCompactMap large = { + {1, 10}, + {2, 20}, + {3, 30}, + }; + copyAssign(&large, &large); + EXPECT_EQ(large.size(), 3u); + EXPECT_EQ(large.at(3), 30); + moveAssign(&large, &large); + EXPECT_EQ(large.size(), 3u); + EXPECT_EQ(large.at(3), 30); +} + +struct TCounted +{ + static int Alive; + + int Value; + + explicit TCounted(int value) + : Value(value) + { + ++Alive; + } + + TCounted(const TCounted& other) + : Value(other.Value) + { + ++Alive; + } + + TCounted(TCounted&& other) noexcept + : Value(other.Value) + { + ++Alive; + } + + TCounted& operator=(const TCounted&) = default; + TCounted& operator=(TCounted&&) = default; + + ~TCounted() + { + --Alive; + } +}; + +int TCounted::Alive = 0; + +TEST(TCompactMapTest, DestroysAllValues) +{ + EXPECT_EQ(TCounted::Alive, 0); + { + TCompactMap map; + map.emplace(1, 10); + map.emplace(2, 20); + map.emplace(3, 30); + + auto copy = map; + auto moved = std::move(copy); + EXPECT_TRUE(copy.empty()); + EXPECT_EQ(moved.size(), 3u); + } + EXPECT_EQ(TCounted::Alive, 0); +} + +//////////////////////////////////////////////////////////////////////////////// +// Move-only value support. +//////////////////////////////////////////////////////////////////////////////// + +struct TMoveOnly +{ + int Value; + + explicit TMoveOnly(int value) + : Value(value) + { } + + TMoveOnly(const TMoveOnly&) = delete; + TMoveOnly& operator=(const TMoveOnly&) = delete; + + TMoveOnly(TMoveOnly&&) = default; + TMoveOnly& operator=(TMoveOnly&&) = default; +}; + +TEST(TCompactMapTest, MoveOnlyValueSmallAndUpgrade) +{ + TCompactMap m; + + auto [it1, inserted1] = m.emplace(1, 10); + ASSERT_TRUE(inserted1); + EXPECT_EQ(it1->second.Value, 10); + + auto [it2, inserted2] = m.emplace(2, 20); + ASSERT_TRUE(inserted2); + EXPECT_EQ(it2->second.Value, 20); + + // try_emplace should move-construct without copying. + auto [it3, inserted3] = m.try_emplace(3, 30); + ASSERT_TRUE(inserted3); + EXPECT_EQ(it3->second.Value, 30); + + // All elements survive the upgrade to std::map. + EXPECT_EQ(m.size(), 3u); + EXPECT_EQ(m.find(1)->second.Value, 10); + EXPECT_EQ(m.find(2)->second.Value, 20); + EXPECT_EQ(m.find(3)->second.Value, 30); +} + +//////////////////////////////////////////////////////////////////////////////// + +} // namespace +} // namespace NYT diff --git a/library/cpp/yt/compact_containers/unittests/compact_vector_ut.cpp b/library/cpp/yt/compact_containers/unittests/compact_vector_ut.cpp index ecc0a745d4a..fd08fb92baf 100644 --- a/library/cpp/yt/compact_containers/unittests/compact_vector_ut.cpp +++ b/library/cpp/yt/compact_containers/unittests/compact_vector_ut.cpp @@ -19,7 +19,10 @@ #include #include +#include #include +#include +#include #include @@ -1091,6 +1094,171 @@ TEST(CompactVectorTest, ZeroPaddingOnHeapMeta) { } } +template +using TTrivialVector = TCompactVector; + +static_assert(sizeof(TTrivialVector<0>) == sizeof(TCompactVectorOnHeapStorage*)); +static_assert(sizeof(TTrivialVector<14>) == 64); +static_assert(sizeof(TTrivialVector<16>) > 64); + +TEST(CompactVectorTest, TrivialMoveAndSwapDoNotDependOnZeroedStorage) { + using TVector = TTrivialVector<4>; + + alignas(TVector) unsigned char lhsStorage[sizeof(TVector)]; + alignas(TVector) unsigned char rhsStorage[sizeof(TVector)]; + std::memset(lhsStorage, 0xa5, sizeof(lhsStorage)); + std::memset(rhsStorage, 0x5a, sizeof(rhsStorage)); + + auto* lhs = ::new (lhsStorage) TVector{1, 2}; + auto* rhs = ::new (rhsStorage) TVector{3, 4, 5, 6, 7}; + + lhs->swap(*rhs); + EXPECT_THAT(*lhs, ::testing::ElementsAre(3, 4, 5, 6, 7)); + EXPECT_THAT(*rhs, ::testing::ElementsAre(1, 2)); + + TVector moved(std::move(*rhs)); + EXPECT_THAT(moved, ::testing::ElementsAre(1, 2)); + EXPECT_TRUE(rhs->empty()); + EXPECT_EQ(4u, rhs->capacity()); + rhs->push_back(42); + EXPECT_THAT(*rhs, ::testing::ElementsAre(42)); + + std::destroy_at(lhs); + std::destroy_at(rhs); +} + +template +void FillTrivialVector(TTrivialVector* vector, size_t size, int firstValue) { + for (size_t index = 0; index < size; ++index) { + vector->push_back(firstValue + index); + } +} + +template +void ExpectMovedFrom(const TTrivialVector& vector) { + EXPECT_TRUE(vector.empty()); + EXPECT_EQ(N, vector.capacity()); +} + +template +void TestTrivialMoveConstruct() { + for (size_t sourceSize : {size_t{0}, N / 2, N + 2}) { + TTrivialVector source; + FillTrivialVector(&source, sourceSize, 10); + std::vector expected(source.begin(), source.end()); + auto* sourceData = source.data(); + + TTrivialVector destination(std::move(source)); + + EXPECT_THAT(destination, ::testing::ElementsAreArray(expected)); + if (sourceSize > N) { + EXPECT_EQ(sourceData, destination.data()); + } + ExpectMovedFrom(source); + source.push_back(42); + EXPECT_THAT(source, ::testing::ElementsAre(42)); + } +} + +TEST(CompactVectorTest, TrivialMoveConstruct) { + TestTrivialMoveConstruct<0>(); + TestTrivialMoveConstruct<4>(); + TestTrivialMoveConstruct<14>(); + TestTrivialMoveConstruct<16>(); +} + +template +void TestTrivialMoveAssign() { + for (size_t sourceSize : {size_t{0}, N / 2, N + 2}) { + for (size_t destinationSize : {size_t{0}, N / 2, N + 2}) { + TTrivialVector source; + TTrivialVector destination; + FillTrivialVector(&source, sourceSize, 10); + FillTrivialVector(&destination, destinationSize, 100); + std::vector expected(source.begin(), source.end()); + auto* sourceData = source.data(); + auto* destinationData = destination.data(); + auto destinationCapacity = destination.capacity(); + + destination = std::move(source); + + EXPECT_THAT(destination, ::testing::ElementsAreArray(expected)); + if (sourceSize > N) { + EXPECT_EQ(sourceData, destination.data()); + } else if (destinationSize > N) { + EXPECT_EQ(destinationData, destination.data()); + EXPECT_EQ(destinationCapacity, destination.capacity()); + } else { + EXPECT_EQ(N, destination.capacity()); + } + ExpectMovedFrom(source); + source.push_back(42); + EXPECT_THAT(source, ::testing::ElementsAre(42)); + } + } + + for (size_t size : {size_t{0}, N / 2, N + 2}) { + TTrivialVector vector; + FillTrivialVector(&vector, size, 10); + std::vector expected(vector.begin(), vector.end()); + auto* data = vector.data(); + + vector = std::move(vector); + + EXPECT_THAT(vector, ::testing::ElementsAreArray(expected)); + EXPECT_EQ(data, vector.data()); + } +} + +TEST(CompactVectorTest, TrivialMoveAssign) { + TestTrivialMoveAssign<4>(); + TestTrivialMoveAssign<14>(); + TestTrivialMoveAssign<16>(); +} + +template +void TestTrivialSwap() { + for (size_t lhsSize : {size_t{0}, size_t{1}, N, N + 2}) { + for (size_t rhsSize : {size_t{0}, size_t{1}, N, N + 2}) { + TTrivialVector lhs; + TTrivialVector rhs; + FillTrivialVector(&lhs, lhsSize, 10); + FillTrivialVector(&rhs, rhsSize, 100); + std::vector expectedLhs(rhs.begin(), rhs.end()); + std::vector expectedRhs(lhs.begin(), lhs.end()); + auto* lhsData = lhs.data(); + auto* rhsData = rhs.data(); + + lhs.swap(rhs); + + EXPECT_THAT(lhs, ::testing::ElementsAreArray(expectedLhs)); + EXPECT_THAT(rhs, ::testing::ElementsAreArray(expectedRhs)); + if (lhsSize > N) { + EXPECT_EQ(lhsData, rhs.data()); + } + if (rhsSize > N) { + EXPECT_EQ(rhsData, lhs.data()); + } + lhs.push_back(42); + rhs.push_back(43); + EXPECT_EQ(42, lhs.back()); + EXPECT_EQ(43, rhs.back()); + } + } +} + +TEST(CompactVectorTest, TrivialSwap) { + TestTrivialSwap<0>(); + TestTrivialSwap<4>(); + TestTrivialSwap<14>(); + TestTrivialSwap<16>(); +} + +TEST(CompactVectorTest, ConstinitDefaultConstruction) { + static constinit TTrivialVector<32> Vector; + EXPECT_TRUE(Vector.empty()); +} + //////////////////////////////////////////////////////////////////////////////// } // namespace diff --git a/library/cpp/yt/memory/intrusive_ptr.h b/library/cpp/yt/memory/intrusive_ptr.h index 613b2f58c49..07b76128514 100644 --- a/library/cpp/yt/memory/intrusive_ptr.h +++ b/library/cpp/yt/memory/intrusive_ptr.h @@ -320,10 +320,10 @@ bool operator==(const TIntrusivePtr& lhs, std::nullptr_t) //////////////////////////////////////////////////////////////////////////////// //! Abseil hash support for TIntrusivePtr. -template -THash AbslHashValue(THash hash, const TIntrusivePtr& ptr) +template +THashState AbslHashValue(THashState hash, const TIntrusivePtr& ptr) { - return THash::combine(std::move(hash), ptr.Get()); + return THashState::combine(std::move(hash), ptr.Get()); } //////////////////////////////////////////////////////////////////////////////// diff --git a/library/cpp/yt/memory/weak_ptr-inl.h b/library/cpp/yt/memory/weak_ptr-inl.h index 0e5ad795774..19d3d199b9b 100644 --- a/library/cpp/yt/memory/weak_ptr-inl.h +++ b/library/cpp/yt/memory/weak_ptr-inl.h @@ -334,10 +334,10 @@ std::size_t TTransparentWeakPtrHasher::operator()(T* ptr) const //////////////////////////////////////////////////////////////////////////////// //! Abseil hash support for TWeakPtr. -template -THash AbslHashValue(THash hash, const TWeakPtr& ptr) +template +THashState AbslHashValue(THashState hash, const TWeakPtr& ptr) { - return THash::combine(std::move(hash), ptr.Get()); + return THashState::combine(std::move(hash), ptr.Get()); } //////////////////////////////////////////////////////////////////////////////// diff --git a/library/cpp/yt/misc/enum-inl.h b/library/cpp/yt/misc/enum-inl.h index dde79f38021..6df52051cfd 100644 --- a/library/cpp/yt/misc/enum-inl.h +++ b/library/cpp/yt/misc/enum-inl.h @@ -477,22 +477,19 @@ constexpr T TEnumTraits::FromString(TStringBuf literal) //////////////////////////////////////////////////////////////////////////////// -template - requires TEnumTraits::IsBitEnum +template constexpr bool Any(E value) noexcept { return ToUnderlying(value) != 0; } -template - requires TEnumTraits::IsBitEnum +template constexpr bool None(E value) noexcept { return ToUnderlying(value) == 0; } -template - requires TEnumTraits::IsBitEnum +template constexpr int PopCount(E value) { return std::popcount(static_cast>(value)); diff --git a/library/cpp/yt/misc/enum.h b/library/cpp/yt/misc/enum.h index 2aa6bee6f77..c0d947d306d 100644 --- a/library/cpp/yt/misc/enum.h +++ b/library/cpp/yt/misc/enum.h @@ -207,14 +207,20 @@ struct TEnumTraits //////////////////////////////////////////////////////////////////////////////// +template +concept CEnum = TEnumTraits::IsEnum; + +template +concept CBitEnum = TEnumTraits::IsBitEnum; + +//////////////////////////////////////////////////////////////////////////////// + //! Returns |true| iff the enumeration value is not bitwise zero. -template - requires TEnumTraits::IsBitEnum +template constexpr bool Any(E value) noexcept; //! Returns |true| iff the enumeration value is bitwise zero. -template - requires TEnumTraits::IsBitEnum +template constexpr bool None(E value) noexcept; //! Returns the number of set bits in |value|. @@ -230,8 +236,7 @@ constexpr bool None(E value) noexcept; //! ); //! //! `PopCount(EMyEnum::Both)` will return 2. -template - requires TEnumTraits::IsBitEnum +template constexpr int PopCount(E value); //////////////////////////////////////////////////////////////////////////////// diff --git a/library/cpp/yt/misc/guid-inl.h b/library/cpp/yt/misc/guid-inl.h index fd5d988f6cc..f688ca8b8ab 100644 --- a/library/cpp/yt/misc/guid-inl.h +++ b/library/cpp/yt/misc/guid-inl.h @@ -74,10 +74,10 @@ Y_FORCE_INLINE std::strong_ordering operator <=> (TGuid lhs, TGuid rhs) noexcept //////////////////////////////////////////////////////////////////////////////// //! Abseil hash support for TGuid. -template -THash AbslHashValue(THash hash, const TGuid& guid) +template +THashState AbslHashValue(THashState hash, const TGuid& guid) { - return THash::combine(std::move(hash), guid.Parts64[0], guid.Parts64[1]); + return THashState::combine(std::move(hash), guid.Parts64[0], guid.Parts64[1]); } //////////////////////////////////////////////////////////////////////////////// diff --git a/library/cpp/yt/misc/no_op.h b/library/cpp/yt/misc/no_op.h new file mode 100644 index 00000000000..8054a995186 --- /dev/null +++ b/library/cpp/yt/misc/no_op.h @@ -0,0 +1,21 @@ +#pragma once + +namespace NYT { + +//////////////////////////////////////////////////////////////////////////////// + +//! A callable that does nothing. +/*! + * Serves as a default template argument where a lambda cannot: each spelling of a + * default argument instantiates its own closure type, so a declaration and its + * definition would disagree on the enclosing type. + */ +struct TNoOp +{ + void operator()() const + { } +}; + +//////////////////////////////////////////////////////////////////////////////// + +} // namespace NYT diff --git a/library/cpp/yt/misc/preprocessor-gen.h b/library/cpp/yt/misc/preprocessor-gen.h index 3e2805a1aae..fc8da12b85f 100644 --- a/library/cpp/yt/misc/preprocessor-gen.h +++ b/library/cpp/yt/misc/preprocessor-gen.h @@ -317,6 +317,7 @@ #define PP_COUNT_CONST_PP_COUNT_IMPL_297 297 #define PP_COUNT_CONST_PP_COUNT_IMPL_298 298 #define PP_COUNT_CONST_PP_COUNT_IMPL_299 299 +#define PP_COUNT_CONST_PP_COUNT_IMPL_300 300 #define PP_COUNT_IMPL_0(_) PP_COUNT_IMPL_1 #define PP_COUNT_IMPL_1(_) PP_COUNT_IMPL_2 #define PP_COUNT_IMPL_2(_) PP_COUNT_IMPL_3 @@ -1236,6 +1237,310 @@ //////////////////////////////////////////////////////////////////////////////// #define PP_TAIL_IMPL(seq) PP_KILL_IMPL(seq, 1) +//////////////////////////////////////////////////////////////////////////////// +#define PP_RANGE_IMPL(first, last) PP_KILL(PP_CONCAT(PP_RANGE_TO_, last), first) +#define PP_RANGE_TO_0 (0) +#define PP_RANGE_TO_1 PP_RANGE_TO_0(1) +#define PP_RANGE_TO_2 PP_RANGE_TO_1(2) +#define PP_RANGE_TO_3 PP_RANGE_TO_2(3) +#define PP_RANGE_TO_4 PP_RANGE_TO_3(4) +#define PP_RANGE_TO_5 PP_RANGE_TO_4(5) +#define PP_RANGE_TO_6 PP_RANGE_TO_5(6) +#define PP_RANGE_TO_7 PP_RANGE_TO_6(7) +#define PP_RANGE_TO_8 PP_RANGE_TO_7(8) +#define PP_RANGE_TO_9 PP_RANGE_TO_8(9) +#define PP_RANGE_TO_10 PP_RANGE_TO_9(10) +#define PP_RANGE_TO_11 PP_RANGE_TO_10(11) +#define PP_RANGE_TO_12 PP_RANGE_TO_11(12) +#define PP_RANGE_TO_13 PP_RANGE_TO_12(13) +#define PP_RANGE_TO_14 PP_RANGE_TO_13(14) +#define PP_RANGE_TO_15 PP_RANGE_TO_14(15) +#define PP_RANGE_TO_16 PP_RANGE_TO_15(16) +#define PP_RANGE_TO_17 PP_RANGE_TO_16(17) +#define PP_RANGE_TO_18 PP_RANGE_TO_17(18) +#define PP_RANGE_TO_19 PP_RANGE_TO_18(19) +#define PP_RANGE_TO_20 PP_RANGE_TO_19(20) +#define PP_RANGE_TO_21 PP_RANGE_TO_20(21) +#define PP_RANGE_TO_22 PP_RANGE_TO_21(22) +#define PP_RANGE_TO_23 PP_RANGE_TO_22(23) +#define PP_RANGE_TO_24 PP_RANGE_TO_23(24) +#define PP_RANGE_TO_25 PP_RANGE_TO_24(25) +#define PP_RANGE_TO_26 PP_RANGE_TO_25(26) +#define PP_RANGE_TO_27 PP_RANGE_TO_26(27) +#define PP_RANGE_TO_28 PP_RANGE_TO_27(28) +#define PP_RANGE_TO_29 PP_RANGE_TO_28(29) +#define PP_RANGE_TO_30 PP_RANGE_TO_29(30) +#define PP_RANGE_TO_31 PP_RANGE_TO_30(31) +#define PP_RANGE_TO_32 PP_RANGE_TO_31(32) +#define PP_RANGE_TO_33 PP_RANGE_TO_32(33) +#define PP_RANGE_TO_34 PP_RANGE_TO_33(34) +#define PP_RANGE_TO_35 PP_RANGE_TO_34(35) +#define PP_RANGE_TO_36 PP_RANGE_TO_35(36) +#define PP_RANGE_TO_37 PP_RANGE_TO_36(37) +#define PP_RANGE_TO_38 PP_RANGE_TO_37(38) +#define PP_RANGE_TO_39 PP_RANGE_TO_38(39) +#define PP_RANGE_TO_40 PP_RANGE_TO_39(40) +#define PP_RANGE_TO_41 PP_RANGE_TO_40(41) +#define PP_RANGE_TO_42 PP_RANGE_TO_41(42) +#define PP_RANGE_TO_43 PP_RANGE_TO_42(43) +#define PP_RANGE_TO_44 PP_RANGE_TO_43(44) +#define PP_RANGE_TO_45 PP_RANGE_TO_44(45) +#define PP_RANGE_TO_46 PP_RANGE_TO_45(46) +#define PP_RANGE_TO_47 PP_RANGE_TO_46(47) +#define PP_RANGE_TO_48 PP_RANGE_TO_47(48) +#define PP_RANGE_TO_49 PP_RANGE_TO_48(49) +#define PP_RANGE_TO_50 PP_RANGE_TO_49(50) +#define PP_RANGE_TO_51 PP_RANGE_TO_50(51) +#define PP_RANGE_TO_52 PP_RANGE_TO_51(52) +#define PP_RANGE_TO_53 PP_RANGE_TO_52(53) +#define PP_RANGE_TO_54 PP_RANGE_TO_53(54) +#define PP_RANGE_TO_55 PP_RANGE_TO_54(55) +#define PP_RANGE_TO_56 PP_RANGE_TO_55(56) +#define PP_RANGE_TO_57 PP_RANGE_TO_56(57) +#define PP_RANGE_TO_58 PP_RANGE_TO_57(58) +#define PP_RANGE_TO_59 PP_RANGE_TO_58(59) +#define PP_RANGE_TO_60 PP_RANGE_TO_59(60) +#define PP_RANGE_TO_61 PP_RANGE_TO_60(61) +#define PP_RANGE_TO_62 PP_RANGE_TO_61(62) +#define PP_RANGE_TO_63 PP_RANGE_TO_62(63) +#define PP_RANGE_TO_64 PP_RANGE_TO_63(64) +#define PP_RANGE_TO_65 PP_RANGE_TO_64(65) +#define PP_RANGE_TO_66 PP_RANGE_TO_65(66) +#define PP_RANGE_TO_67 PP_RANGE_TO_66(67) +#define PP_RANGE_TO_68 PP_RANGE_TO_67(68) +#define PP_RANGE_TO_69 PP_RANGE_TO_68(69) +#define PP_RANGE_TO_70 PP_RANGE_TO_69(70) +#define PP_RANGE_TO_71 PP_RANGE_TO_70(71) +#define PP_RANGE_TO_72 PP_RANGE_TO_71(72) +#define PP_RANGE_TO_73 PP_RANGE_TO_72(73) +#define PP_RANGE_TO_74 PP_RANGE_TO_73(74) +#define PP_RANGE_TO_75 PP_RANGE_TO_74(75) +#define PP_RANGE_TO_76 PP_RANGE_TO_75(76) +#define PP_RANGE_TO_77 PP_RANGE_TO_76(77) +#define PP_RANGE_TO_78 PP_RANGE_TO_77(78) +#define PP_RANGE_TO_79 PP_RANGE_TO_78(79) +#define PP_RANGE_TO_80 PP_RANGE_TO_79(80) +#define PP_RANGE_TO_81 PP_RANGE_TO_80(81) +#define PP_RANGE_TO_82 PP_RANGE_TO_81(82) +#define PP_RANGE_TO_83 PP_RANGE_TO_82(83) +#define PP_RANGE_TO_84 PP_RANGE_TO_83(84) +#define PP_RANGE_TO_85 PP_RANGE_TO_84(85) +#define PP_RANGE_TO_86 PP_RANGE_TO_85(86) +#define PP_RANGE_TO_87 PP_RANGE_TO_86(87) +#define PP_RANGE_TO_88 PP_RANGE_TO_87(88) +#define PP_RANGE_TO_89 PP_RANGE_TO_88(89) +#define PP_RANGE_TO_90 PP_RANGE_TO_89(90) +#define PP_RANGE_TO_91 PP_RANGE_TO_90(91) +#define PP_RANGE_TO_92 PP_RANGE_TO_91(92) +#define PP_RANGE_TO_93 PP_RANGE_TO_92(93) +#define PP_RANGE_TO_94 PP_RANGE_TO_93(94) +#define PP_RANGE_TO_95 PP_RANGE_TO_94(95) +#define PP_RANGE_TO_96 PP_RANGE_TO_95(96) +#define PP_RANGE_TO_97 PP_RANGE_TO_96(97) +#define PP_RANGE_TO_98 PP_RANGE_TO_97(98) +#define PP_RANGE_TO_99 PP_RANGE_TO_98(99) +#define PP_RANGE_TO_100 PP_RANGE_TO_99(100) +#define PP_RANGE_TO_101 PP_RANGE_TO_100(101) +#define PP_RANGE_TO_102 PP_RANGE_TO_101(102) +#define PP_RANGE_TO_103 PP_RANGE_TO_102(103) +#define PP_RANGE_TO_104 PP_RANGE_TO_103(104) +#define PP_RANGE_TO_105 PP_RANGE_TO_104(105) +#define PP_RANGE_TO_106 PP_RANGE_TO_105(106) +#define PP_RANGE_TO_107 PP_RANGE_TO_106(107) +#define PP_RANGE_TO_108 PP_RANGE_TO_107(108) +#define PP_RANGE_TO_109 PP_RANGE_TO_108(109) +#define PP_RANGE_TO_110 PP_RANGE_TO_109(110) +#define PP_RANGE_TO_111 PP_RANGE_TO_110(111) +#define PP_RANGE_TO_112 PP_RANGE_TO_111(112) +#define PP_RANGE_TO_113 PP_RANGE_TO_112(113) +#define PP_RANGE_TO_114 PP_RANGE_TO_113(114) +#define PP_RANGE_TO_115 PP_RANGE_TO_114(115) +#define PP_RANGE_TO_116 PP_RANGE_TO_115(116) +#define PP_RANGE_TO_117 PP_RANGE_TO_116(117) +#define PP_RANGE_TO_118 PP_RANGE_TO_117(118) +#define PP_RANGE_TO_119 PP_RANGE_TO_118(119) +#define PP_RANGE_TO_120 PP_RANGE_TO_119(120) +#define PP_RANGE_TO_121 PP_RANGE_TO_120(121) +#define PP_RANGE_TO_122 PP_RANGE_TO_121(122) +#define PP_RANGE_TO_123 PP_RANGE_TO_122(123) +#define PP_RANGE_TO_124 PP_RANGE_TO_123(124) +#define PP_RANGE_TO_125 PP_RANGE_TO_124(125) +#define PP_RANGE_TO_126 PP_RANGE_TO_125(126) +#define PP_RANGE_TO_127 PP_RANGE_TO_126(127) +#define PP_RANGE_TO_128 PP_RANGE_TO_127(128) +#define PP_RANGE_TO_129 PP_RANGE_TO_128(129) +#define PP_RANGE_TO_130 PP_RANGE_TO_129(130) +#define PP_RANGE_TO_131 PP_RANGE_TO_130(131) +#define PP_RANGE_TO_132 PP_RANGE_TO_131(132) +#define PP_RANGE_TO_133 PP_RANGE_TO_132(133) +#define PP_RANGE_TO_134 PP_RANGE_TO_133(134) +#define PP_RANGE_TO_135 PP_RANGE_TO_134(135) +#define PP_RANGE_TO_136 PP_RANGE_TO_135(136) +#define PP_RANGE_TO_137 PP_RANGE_TO_136(137) +#define PP_RANGE_TO_138 PP_RANGE_TO_137(138) +#define PP_RANGE_TO_139 PP_RANGE_TO_138(139) +#define PP_RANGE_TO_140 PP_RANGE_TO_139(140) +#define PP_RANGE_TO_141 PP_RANGE_TO_140(141) +#define PP_RANGE_TO_142 PP_RANGE_TO_141(142) +#define PP_RANGE_TO_143 PP_RANGE_TO_142(143) +#define PP_RANGE_TO_144 PP_RANGE_TO_143(144) +#define PP_RANGE_TO_145 PP_RANGE_TO_144(145) +#define PP_RANGE_TO_146 PP_RANGE_TO_145(146) +#define PP_RANGE_TO_147 PP_RANGE_TO_146(147) +#define PP_RANGE_TO_148 PP_RANGE_TO_147(148) +#define PP_RANGE_TO_149 PP_RANGE_TO_148(149) +#define PP_RANGE_TO_150 PP_RANGE_TO_149(150) +#define PP_RANGE_TO_151 PP_RANGE_TO_150(151) +#define PP_RANGE_TO_152 PP_RANGE_TO_151(152) +#define PP_RANGE_TO_153 PP_RANGE_TO_152(153) +#define PP_RANGE_TO_154 PP_RANGE_TO_153(154) +#define PP_RANGE_TO_155 PP_RANGE_TO_154(155) +#define PP_RANGE_TO_156 PP_RANGE_TO_155(156) +#define PP_RANGE_TO_157 PP_RANGE_TO_156(157) +#define PP_RANGE_TO_158 PP_RANGE_TO_157(158) +#define PP_RANGE_TO_159 PP_RANGE_TO_158(159) +#define PP_RANGE_TO_160 PP_RANGE_TO_159(160) +#define PP_RANGE_TO_161 PP_RANGE_TO_160(161) +#define PP_RANGE_TO_162 PP_RANGE_TO_161(162) +#define PP_RANGE_TO_163 PP_RANGE_TO_162(163) +#define PP_RANGE_TO_164 PP_RANGE_TO_163(164) +#define PP_RANGE_TO_165 PP_RANGE_TO_164(165) +#define PP_RANGE_TO_166 PP_RANGE_TO_165(166) +#define PP_RANGE_TO_167 PP_RANGE_TO_166(167) +#define PP_RANGE_TO_168 PP_RANGE_TO_167(168) +#define PP_RANGE_TO_169 PP_RANGE_TO_168(169) +#define PP_RANGE_TO_170 PP_RANGE_TO_169(170) +#define PP_RANGE_TO_171 PP_RANGE_TO_170(171) +#define PP_RANGE_TO_172 PP_RANGE_TO_171(172) +#define PP_RANGE_TO_173 PP_RANGE_TO_172(173) +#define PP_RANGE_TO_174 PP_RANGE_TO_173(174) +#define PP_RANGE_TO_175 PP_RANGE_TO_174(175) +#define PP_RANGE_TO_176 PP_RANGE_TO_175(176) +#define PP_RANGE_TO_177 PP_RANGE_TO_176(177) +#define PP_RANGE_TO_178 PP_RANGE_TO_177(178) +#define PP_RANGE_TO_179 PP_RANGE_TO_178(179) +#define PP_RANGE_TO_180 PP_RANGE_TO_179(180) +#define PP_RANGE_TO_181 PP_RANGE_TO_180(181) +#define PP_RANGE_TO_182 PP_RANGE_TO_181(182) +#define PP_RANGE_TO_183 PP_RANGE_TO_182(183) +#define PP_RANGE_TO_184 PP_RANGE_TO_183(184) +#define PP_RANGE_TO_185 PP_RANGE_TO_184(185) +#define PP_RANGE_TO_186 PP_RANGE_TO_185(186) +#define PP_RANGE_TO_187 PP_RANGE_TO_186(187) +#define PP_RANGE_TO_188 PP_RANGE_TO_187(188) +#define PP_RANGE_TO_189 PP_RANGE_TO_188(189) +#define PP_RANGE_TO_190 PP_RANGE_TO_189(190) +#define PP_RANGE_TO_191 PP_RANGE_TO_190(191) +#define PP_RANGE_TO_192 PP_RANGE_TO_191(192) +#define PP_RANGE_TO_193 PP_RANGE_TO_192(193) +#define PP_RANGE_TO_194 PP_RANGE_TO_193(194) +#define PP_RANGE_TO_195 PP_RANGE_TO_194(195) +#define PP_RANGE_TO_196 PP_RANGE_TO_195(196) +#define PP_RANGE_TO_197 PP_RANGE_TO_196(197) +#define PP_RANGE_TO_198 PP_RANGE_TO_197(198) +#define PP_RANGE_TO_199 PP_RANGE_TO_198(199) +#define PP_RANGE_TO_200 PP_RANGE_TO_199(200) +#define PP_RANGE_TO_201 PP_RANGE_TO_200(201) +#define PP_RANGE_TO_202 PP_RANGE_TO_201(202) +#define PP_RANGE_TO_203 PP_RANGE_TO_202(203) +#define PP_RANGE_TO_204 PP_RANGE_TO_203(204) +#define PP_RANGE_TO_205 PP_RANGE_TO_204(205) +#define PP_RANGE_TO_206 PP_RANGE_TO_205(206) +#define PP_RANGE_TO_207 PP_RANGE_TO_206(207) +#define PP_RANGE_TO_208 PP_RANGE_TO_207(208) +#define PP_RANGE_TO_209 PP_RANGE_TO_208(209) +#define PP_RANGE_TO_210 PP_RANGE_TO_209(210) +#define PP_RANGE_TO_211 PP_RANGE_TO_210(211) +#define PP_RANGE_TO_212 PP_RANGE_TO_211(212) +#define PP_RANGE_TO_213 PP_RANGE_TO_212(213) +#define PP_RANGE_TO_214 PP_RANGE_TO_213(214) +#define PP_RANGE_TO_215 PP_RANGE_TO_214(215) +#define PP_RANGE_TO_216 PP_RANGE_TO_215(216) +#define PP_RANGE_TO_217 PP_RANGE_TO_216(217) +#define PP_RANGE_TO_218 PP_RANGE_TO_217(218) +#define PP_RANGE_TO_219 PP_RANGE_TO_218(219) +#define PP_RANGE_TO_220 PP_RANGE_TO_219(220) +#define PP_RANGE_TO_221 PP_RANGE_TO_220(221) +#define PP_RANGE_TO_222 PP_RANGE_TO_221(222) +#define PP_RANGE_TO_223 PP_RANGE_TO_222(223) +#define PP_RANGE_TO_224 PP_RANGE_TO_223(224) +#define PP_RANGE_TO_225 PP_RANGE_TO_224(225) +#define PP_RANGE_TO_226 PP_RANGE_TO_225(226) +#define PP_RANGE_TO_227 PP_RANGE_TO_226(227) +#define PP_RANGE_TO_228 PP_RANGE_TO_227(228) +#define PP_RANGE_TO_229 PP_RANGE_TO_228(229) +#define PP_RANGE_TO_230 PP_RANGE_TO_229(230) +#define PP_RANGE_TO_231 PP_RANGE_TO_230(231) +#define PP_RANGE_TO_232 PP_RANGE_TO_231(232) +#define PP_RANGE_TO_233 PP_RANGE_TO_232(233) +#define PP_RANGE_TO_234 PP_RANGE_TO_233(234) +#define PP_RANGE_TO_235 PP_RANGE_TO_234(235) +#define PP_RANGE_TO_236 PP_RANGE_TO_235(236) +#define PP_RANGE_TO_237 PP_RANGE_TO_236(237) +#define PP_RANGE_TO_238 PP_RANGE_TO_237(238) +#define PP_RANGE_TO_239 PP_RANGE_TO_238(239) +#define PP_RANGE_TO_240 PP_RANGE_TO_239(240) +#define PP_RANGE_TO_241 PP_RANGE_TO_240(241) +#define PP_RANGE_TO_242 PP_RANGE_TO_241(242) +#define PP_RANGE_TO_243 PP_RANGE_TO_242(243) +#define PP_RANGE_TO_244 PP_RANGE_TO_243(244) +#define PP_RANGE_TO_245 PP_RANGE_TO_244(245) +#define PP_RANGE_TO_246 PP_RANGE_TO_245(246) +#define PP_RANGE_TO_247 PP_RANGE_TO_246(247) +#define PP_RANGE_TO_248 PP_RANGE_TO_247(248) +#define PP_RANGE_TO_249 PP_RANGE_TO_248(249) +#define PP_RANGE_TO_250 PP_RANGE_TO_249(250) +#define PP_RANGE_TO_251 PP_RANGE_TO_250(251) +#define PP_RANGE_TO_252 PP_RANGE_TO_251(252) +#define PP_RANGE_TO_253 PP_RANGE_TO_252(253) +#define PP_RANGE_TO_254 PP_RANGE_TO_253(254) +#define PP_RANGE_TO_255 PP_RANGE_TO_254(255) +#define PP_RANGE_TO_256 PP_RANGE_TO_255(256) +#define PP_RANGE_TO_257 PP_RANGE_TO_256(257) +#define PP_RANGE_TO_258 PP_RANGE_TO_257(258) +#define PP_RANGE_TO_259 PP_RANGE_TO_258(259) +#define PP_RANGE_TO_260 PP_RANGE_TO_259(260) +#define PP_RANGE_TO_261 PP_RANGE_TO_260(261) +#define PP_RANGE_TO_262 PP_RANGE_TO_261(262) +#define PP_RANGE_TO_263 PP_RANGE_TO_262(263) +#define PP_RANGE_TO_264 PP_RANGE_TO_263(264) +#define PP_RANGE_TO_265 PP_RANGE_TO_264(265) +#define PP_RANGE_TO_266 PP_RANGE_TO_265(266) +#define PP_RANGE_TO_267 PP_RANGE_TO_266(267) +#define PP_RANGE_TO_268 PP_RANGE_TO_267(268) +#define PP_RANGE_TO_269 PP_RANGE_TO_268(269) +#define PP_RANGE_TO_270 PP_RANGE_TO_269(270) +#define PP_RANGE_TO_271 PP_RANGE_TO_270(271) +#define PP_RANGE_TO_272 PP_RANGE_TO_271(272) +#define PP_RANGE_TO_273 PP_RANGE_TO_272(273) +#define PP_RANGE_TO_274 PP_RANGE_TO_273(274) +#define PP_RANGE_TO_275 PP_RANGE_TO_274(275) +#define PP_RANGE_TO_276 PP_RANGE_TO_275(276) +#define PP_RANGE_TO_277 PP_RANGE_TO_276(277) +#define PP_RANGE_TO_278 PP_RANGE_TO_277(278) +#define PP_RANGE_TO_279 PP_RANGE_TO_278(279) +#define PP_RANGE_TO_280 PP_RANGE_TO_279(280) +#define PP_RANGE_TO_281 PP_RANGE_TO_280(281) +#define PP_RANGE_TO_282 PP_RANGE_TO_281(282) +#define PP_RANGE_TO_283 PP_RANGE_TO_282(283) +#define PP_RANGE_TO_284 PP_RANGE_TO_283(284) +#define PP_RANGE_TO_285 PP_RANGE_TO_284(285) +#define PP_RANGE_TO_286 PP_RANGE_TO_285(286) +#define PP_RANGE_TO_287 PP_RANGE_TO_286(287) +#define PP_RANGE_TO_288 PP_RANGE_TO_287(288) +#define PP_RANGE_TO_289 PP_RANGE_TO_288(289) +#define PP_RANGE_TO_290 PP_RANGE_TO_289(290) +#define PP_RANGE_TO_291 PP_RANGE_TO_290(291) +#define PP_RANGE_TO_292 PP_RANGE_TO_291(292) +#define PP_RANGE_TO_293 PP_RANGE_TO_292(293) +#define PP_RANGE_TO_294 PP_RANGE_TO_293(294) +#define PP_RANGE_TO_295 PP_RANGE_TO_294(295) +#define PP_RANGE_TO_296 PP_RANGE_TO_295(296) +#define PP_RANGE_TO_297 PP_RANGE_TO_296(297) +#define PP_RANGE_TO_298 PP_RANGE_TO_297(298) +#define PP_RANGE_TO_299 PP_RANGE_TO_298(299) +#define PP_RANGE_TO_300 PP_RANGE_TO_299(300) + //////////////////////////////////////////////////////////////////////////////// #define PP_FOR_EACH_IMPL(what, seq) PP_CONCAT(PP_FOR_EACH_IMPL_, \ PP_COUNT(seq))(what, seq) diff --git a/library/cpp/yt/misc/preprocessor.h b/library/cpp/yt/misc/preprocessor.h index 818c3be3546..7b27a817048 100644 --- a/library/cpp/yt/misc/preprocessor.h +++ b/library/cpp/yt/misc/preprocessor.h @@ -104,6 +104,12 @@ */ #define PP_ELEMENT(seq, index) PP_ELEMENT_IMPL(seq, index) +//! Generates a sequence of integers from first to last, inclusive. +/*! The bounds must expand to integer literals satisfying 0 <= first <= last <= 300. + * For example, \code PP_RANGE(1, 3) == (1)(2)(3) \endcode + */ +#define PP_RANGE(first, last) PP_RANGE_IMPL(first, last) + //! Applies the macro to every member of the sequence. /*! For example, * \code diff --git a/library/cpp/yt/misc/strong_typedef-inl.h b/library/cpp/yt/misc/strong_typedef-inl.h index 068990cbe79..a3fdabc1615 100644 --- a/library/cpp/yt/misc/strong_typedef-inl.h +++ b/library/cpp/yt/misc/strong_typedef-inl.h @@ -166,10 +166,10 @@ struct TBasicWrapperTraits> //////////////////////////////////////////////////////////////////////////////// //! Abseil hash support for TStrongTypedef. -template -THash AbslHashValue(THash hash, const TStrongTypedef& value) +template +THashState AbslHashValue(THashState hash, const TStrongTypedef& value) { - return THash::combine(std::move(hash), value.Underlying()); + return THashState::combine(std::move(hash), value.Underlying()); } //////////////////////////////////////////////////////////////////////////////// diff --git a/library/cpp/yt/misc/unittests/preprocessor_ut.cpp b/library/cpp/yt/misc/unittests/preprocessor_ut.cpp index e4c88a7fe8e..0d0fa2d4bec 100644 --- a/library/cpp/yt/misc/unittests/preprocessor_ut.cpp +++ b/library/cpp/yt/misc/unittests/preprocessor_ut.cpp @@ -54,6 +54,39 @@ TEST(TPreprocessorTest, Tail) EXPECT_STREQ("PP_NIL (1)(2)", PP_STRINGIZE(PP_NIL PP_TAIL((0)(1)(2)))); } +TEST(TPreprocessorTest, Range) +{ + EXPECT_EQ(0, PP_HEAD(PP_RANGE(0, 0))); + EXPECT_EQ(1, PP_COUNT(PP_RANGE(0, 0))); + EXPECT_EQ(1, PP_COUNT(PP_RANGE(8, 8))); + EXPECT_EQ(8, PP_HEAD(PP_RANGE(8, 8))); + + EXPECT_EQ(8, PP_COUNT(PP_RANGE(1, 8))); + EXPECT_EQ(1, PP_ELEMENT(PP_RANGE(1, 8), 0)); + EXPECT_EQ(4, PP_ELEMENT(PP_RANGE(1, 8), 3)); + EXPECT_EQ(8, PP_ELEMENT(PP_RANGE(1, 8), 7)); + + EXPECT_EQ(300, PP_COUNT(PP_RANGE(0, 299))); + EXPECT_EQ(300, PP_COUNT(PP_RANGE(1, 300))); + EXPECT_EQ(300, PP_ELEMENT(PP_RANGE(1, 300), 299)); + EXPECT_EQ(1, PP_COUNT(PP_RANGE(300, 300))); + EXPECT_EQ(300, PP_HEAD(PP_RANGE(300, 300))); + +#define TEST_RANGE_FIRST 2 +#define TEST_RANGE_LAST 5 + EXPECT_EQ(4, PP_COUNT(PP_RANGE(TEST_RANGE_FIRST, TEST_RANGE_LAST))); + EXPECT_EQ(2, PP_HEAD(PP_RANGE(TEST_RANGE_FIRST, TEST_RANGE_LAST))); + EXPECT_EQ(5, PP_ELEMENT(PP_RANGE(TEST_RANGE_FIRST, TEST_RANGE_LAST), 3)); +#undef TEST_RANGE_FIRST +#undef TEST_RANGE_LAST + +#define ADD_RANGE_VALUE(value) + value + EXPECT_EQ(36, (0 PP_FOR_EACH(ADD_RANGE_VALUE, PP_RANGE(1, 8)))); + EXPECT_EQ(7, (0 PP_FOR_EACH(ADD_RANGE_VALUE, PP_RANGE(7, 7)))); + EXPECT_EQ(45150, (0 PP_FOR_EACH(ADD_RANGE_VALUE, PP_RANGE(1, 300)))); +#undef ADD_RANGE_VALUE +} + TEST(TPreprocessorTest, ForEach) { EXPECT_STREQ( diff --git a/library/cpp/yt/string/format-inl.h b/library/cpp/yt/string/format-inl.h index e0325754c33..a34569c9be7 100644 --- a/library/cpp/yt/string/format-inl.h +++ b/library/cpp/yt/string/format-inl.h @@ -603,7 +603,7 @@ inline void FormatValue(TStringBuilderBase* builder, const std::string_view& val // std::filesystem::path inline void FormatValue(TStringBuilderBase* builder, const std::filesystem::path& value, TStringBuf spec) { - FormatValue(builder, std::string(value), spec); + FormatValue(builder, value.string(), spec); } #endif diff --git a/library/cpp/yt/string/readme.md b/library/cpp/yt/string/readme.md index 1e19eddb53c..5a5330b04b2 100644 --- a/library/cpp/yt/string/readme.md +++ b/library/cpp/yt/string/readme.md @@ -132,7 +132,7 @@ ROOT/library/cpp/yt/string/format_string-inl.h:36:10: note: in call to '&[] { ... ``` -First line contains the source location where the error occured. Second line contains the function name `CrashCompilerClassIsNotFormattable` which name is the error and template argument is the errorneos type. There are some more lines which would contain incomprehensible garbage --- don't bother reading it. Other compiler errors generated by static analyser (see below) follow the same structure. +First line contains the source location where the error occurred. Second line contains the function name `CrashCompilerClassIsNotFormattable` which name is the error and template argument is the errorneos type. There are some more lines which would contain incomprehensible garbage --- don't bother reading it. Other compiler errors generated by static analyser (see below) follow the same structure. In order to support printing custom type, one must create an overload of `FormatValue` function. If everything is done correctly, concept `CFormattable` should be satisfied and the value printed accordingly. @@ -237,7 +237,7 @@ struct NYT::TFormatArg }; ``` -Now we are able to print the value as format analyser is aware of the new flag 'k'. If we wanted to, we could remove the rest of the default specifiers provided by `TFormatArgBase`, since most of them might not make any sence for your type. +Now we are able to print the value as format analyser is aware of the new flag 'k'. If we wanted to, we could remove the rest of the default specifiers provided by `TFormatArgBase`, since most of them might not make any sense for your type. You can use `TFormatArg` + `FormatValue` to fully support format decorators: ```cpp diff --git a/library/cpp/yt/yson_string/string-inl.h b/library/cpp/yt/yson_string/string-inl.h index bc8d3ed0c63..cb8faa21587 100644 --- a/library/cpp/yt/yson_string/string-inl.h +++ b/library/cpp/yt/yson_string/string-inl.h @@ -100,17 +100,17 @@ inline bool operator != (const TYsonStringBuf& lhs, const TYsonStringBuf& rhs) //////////////////////////////////////////////////////////////////////////////// //! Abseil hash support for TYsonString. -template -THash AbslHashValue(THash hash, const TYsonString& str) +template +THashState AbslHashValue(THashState hash, const TYsonString& str) { - return THash::combine(std::move(hash), str.AsStringBuf()); + return THashState::combine(std::move(hash), str ? str.AsStringBuf() : TStringBuf()); } //! Abseil hash support for TYsonStringBuf. -template -THash AbslHashValue(THash hash, const TYsonStringBuf& str) +template +THashState AbslHashValue(THashState hash, const TYsonStringBuf& str) { - return THash::combine(std::move(hash), str.AsStringBuf()); + return THashState::combine(std::move(hash), str ? str.AsStringBuf() : TStringBuf()); } //////////////////////////////////////////////////////////////////////////////// diff --git a/src/api/grpc/ydb_topic_v1.proto b/src/api/grpc/ydb_topic_v1.proto index 8480125ac1b..0d3c4b379c7 100644 --- a/src/api/grpc/ydb_topic_v1.proto +++ b/src/api/grpc/ydb_topic_v1.proto @@ -115,6 +115,11 @@ service TopicService { // Single commit offset request. rpc CommitOffset(CommitOffsetRequest) returns (CommitOffsetResponse); + // Reset committed offsets for a consumer on all topic partitions (including inactive). + // Each partition is updated independently (not atomic across partitions). + // Drops any active read session for this consumer. + rpc ResetOffset(ResetOffsetRequest) returns (ResetOffsetResponse); + // Add information about offset ranges to the transaction. rpc UpdateOffsetsInTransaction(UpdateOffsetsInTransactionRequest) returns (UpdateOffsetsInTransactionResponse); diff --git a/src/api/grpc/ydb_udf_v1.proto b/src/api/grpc/ydb_udf_v1.proto new file mode 100644 index 00000000000..b860d10ff03 --- /dev/null +++ b/src/api/grpc/ydb_udf_v1.proto @@ -0,0 +1,17 @@ +syntax = "proto3"; + +package Ydb.Udf.V1; +option java_package = "tech.ydb.udf.v1"; + +import "src/api/protos/ydb_udf.proto"; + +service UdfService { + // Bidirectional stream on the wire (YDB streaming infra). Client sends + // metadata + data chunks; server replies with a single UploadModuleResponse. + rpc UploadModule(stream Udf.UploadModuleChunk) returns (stream Udf.UploadModuleResponse); + + rpc DeleteModule(Udf.DeleteModuleRequest) returns (Udf.DeleteModuleResponse); + + rpc ListModules(Udf.ListModulesRequest) returns (Udf.ListModulesResponse); + rpc DescribeModule(Udf.DescribeModuleRequest) returns (Udf.DescribeModuleResponse); +} diff --git a/src/api/protos/draft/ydb_distributed_storage.proto b/src/api/protos/draft/ydb_distributed_storage.proto index fe543b20484..050553fa546 100644 --- a/src/api/protos/draft/ydb_distributed_storage.proto +++ b/src/api/protos/draft/ydb_distributed_storage.proto @@ -461,6 +461,10 @@ message SafetyOptions { bool ignore_degraded_groups = 1; // Skip the group failure-model check. The operation may cause data loss. bool ignore_group_failure_model = 2; + // Allow VDisk reassignment to leave an incorrect physical group layout, even + // if the layout was already incorrect. Target eligibility and space checks + // still apply. Other operations do not use this option. + bool ignore_group_layout_checks = 3; } message ReassignVDiskOptions { diff --git a/src/api/protos/ydb_table.proto b/src/api/protos/ydb_table.proto index 6f3cda7065e..5a144990cca 100644 --- a/src/api/protos/ydb_table.proto +++ b/src/api/protos/ydb_table.proto @@ -80,6 +80,8 @@ message VectorIndexSettings { enum VectorType { VECTOR_TYPE_UNSPECIFIED = 0; VECTOR_TYPE_FLOAT = 1; + VECTOR_TYPE_FLOAT16 = 5; + VECTOR_TYPE_BFLOAT16 = 6; VECTOR_TYPE_UINT8 = 2; VECTOR_TYPE_INT8 = 3; VECTOR_TYPE_BIT = 4; @@ -163,8 +165,8 @@ message FulltextIndexSettings { // See Tokenizer enum optional Tokenizer tokenizer = 1; - // Language used for language-sensitive operations like stopword filtering and stemming - // Example: language = "english" + // Comma-separated languages used for language-sensitive operations like stopword filtering and stemming + // Example: language = "english,russian" // By default is not specified and no language-specific logic is applied optional string language = 2; @@ -220,8 +222,8 @@ message FulltextIndexSettings { optional int32 filter_length_max = 132 [(Ydb.value) = ">= 0"]; // Whether to apply snowball stemming to each token - // Must be used with language option - // Example: language = "english" + // Must be used with language option; languages in a list must use distinct scripts so each word can be routed to one stemmer + // Example: language = "english,russian" // Tokens: ["cars", "beautifully", "conspired"] // Output: ["car", "beauti", "conspir"] optional bool use_filter_snowball = 140; @@ -831,6 +833,7 @@ message ColumnMeta { message EvictionToExternalStorageSettings { // Path to external data source string storage = 1; + optional string object_key_prefix = 2; } message DateTypeColumnModeSettings { diff --git a/src/api/protos/ydb_topic.proto b/src/api/protos/ydb_topic.proto index f9ce1e4e797..6a1df69bd9d 100644 --- a/src/api/protos/ydb_topic.proto +++ b/src/api/protos/ydb_topic.proto @@ -802,6 +802,57 @@ message CommitOffsetResult { } +//////////////////////////////////////////////////////////////////////////////////////////////////////////////////////// +// ResetOffset + + +// Reset committed offsets for a consumer on all topic partitions +// (including inactive partitions after split/merge). +// +// Each partition is updated independently: a successful overall status does not +// mean every partition was reset, and a failure may still leave some partitions +// already rewritten. The request is not atomic across partitions. +// +// The partition applies the new offset with an empty read session id, which +// drops any active read session for this consumer on that partition. +message ResetOffsetRequest { + Ydb.Operations.OperationParams operation_params = 1; + + // Topic path. + string path = 2; + // Consumer name. + string consumer = 3; + + message Earliest { + } + + message Latest { + } + + message FromWrittenAt { + // Offset of the first message with write timestamp >= written_at. + // If none, the partition end offset is used. + google.protobuf.Timestamp written_at = 1; + } + + oneof position { + Earliest earliest = 4; + Latest latest = 5; + FromWrittenAt from_written_at = 6; + } +} + +// Reset offset response sent from server to client. +message ResetOffsetResponse { + // Result of request will be inside operation. + Ydb.Operations.Operation operation = 1; +} + +// Reset offset result message inside ResetOffsetResponse.operation. +message ResetOffsetResult { +} + + //////////////////////////////////////////////////////////////////////////////////////////////////////////////////////// // Control messages diff --git a/src/api/protos/ydb_udf.proto b/src/api/protos/ydb_udf.proto new file mode 100644 index 00000000000..bce0eacb79e --- /dev/null +++ b/src/api/protos/ydb_udf.proto @@ -0,0 +1,154 @@ +syntax = "proto3"; +option cc_enable_arenas = true; + +package Ydb.Udf; +option java_package = "tech.ydb.udf"; + +import "google/protobuf/timestamp.proto"; +import "src/api/protos/annotations/validation.proto"; +import "src/api/protos/ydb_operation.proto"; + +enum ModuleType { + MODULE_TYPE_UNSPECIFIED = 0; + MODULE = 1; + LIBRARY = 2; +} + +enum ModuleKind { + MODULE_KIND_UNSPECIFIED = 0; + WASM = 1; + NATIVE = 2; // Recognized, but not supported yet. +} + +enum CompileStatus { + COMPILE_STATUS_UNSPECIFIED = 0; + PENDING = 1; + COMPILING = 2; + READY = 3; + FAILED = 4; +} + +enum WriteMode { + WRITE_MODE_UNSPECIFIED = 0; // treated as CREATE_OR_REPLACE + CREATE_OR_REPLACE = 1; + CREATE_ONLY = 2; + REPLACE_ONLY = 3; +} + +message ModuleInfo { + reserved 7, 8, 10; + reserved "compile_status", "compile_error", "compile_finished_at"; + + string name = 1; + ModuleType module_type = 2; + ModuleKind module_kind = 11; + string uid = 3; + string md5 = 4; + uint64 size = 5; + uint64 version = 6; + google.protobuf.Timestamp created_at = 9; +} + +message PlatformCompileStatus { + string cpu_spec = 1; + CompileStatus status = 2; + string compile_error = 3; + google.protobuf.Timestamp compile_started_at = 4; + google.protobuf.Timestamp compile_finished_at = 5; +} + +// Params for upload (no body — bytes go in subsequent stream chunks). +message UploadModuleParams { + Ydb.Operations.OperationParams operation_params = 1; + + string manifest_json = 4; // Required; sole source of name, type and kind. + + WriteMode write_mode = 5; + string expected_uid = 6; + + uint64 version = 7; + string expected_md5 = 8; +} + +// Client→server stream messages. Wire form is bidirectional (YDB streaming infra), +// but the server sends exactly one UploadModuleResponse and finishes. +message UploadModuleChunk { + oneof payload { + UploadModuleMetadata metadata = 1; // must be first message + bytes data = 2; // subsequent chunks + } +} + +message UploadModuleMetadata { + UploadModuleParams params = 1; + // Required: the transport reports a client half-close and a broken + // connection alike, so this is the only thing that tells a whole body from + // one cut short. The server refuses the upload unless the bytes it received + // add up to exactly this. + uint64 total_size = 2; +} + +message UploadModuleResponse { + Ydb.Operations.Operation operation = 1; +} + +message UploadModuleResult { + reserved 5; + reserved "compile_status"; + + string name = 1; + string uid = 2; + string md5 = 3; + uint64 size = 4; + bool replaced_existing = 6; +} + +message DeleteModuleRequest { + Ydb.Operations.OperationParams operation_params = 1; + string name = 2 [(required) = true]; + ModuleType module_type = 3; // optional assert + ModuleKind module_kind = 5; // optional assert + string expected_uid = 4; +} + +message DeleteModuleResponse { + Ydb.Operations.Operation operation = 1; +} + +message DeleteModuleResult { +} + +message ListModulesRequest { + reserved 3; + reserved "status_filter"; + + Ydb.Operations.OperationParams operation_params = 1; + ModuleType type_filter = 2; + ModuleKind kind_filter = 6; + uint32 page_size = 4; + string page_token = 5; +} + +message ListModulesResponse { + Ydb.Operations.Operation operation = 1; +} + +message ListModulesResult { + repeated ModuleInfo modules = 1; + string next_page_token = 2; +} + +message DescribeModuleRequest { + Ydb.Operations.OperationParams operation_params = 1; + string name = 2 [(required) = true]; +} + +message DescribeModuleResponse { + Ydb.Operations.Operation operation = 1; +} + +message DescribeModuleResult { + ModuleInfo module_info = 1 [json_name = "module"]; + string manifest_json = 2; + repeated PlatformCompileStatus platforms = 3; +} diff --git a/src/client/impl/internal/grpc_connections/grpc_connections.h b/src/client/impl/internal/grpc_connections/grpc_connections.h index 58901e3a65e..f29ef2a77ff 100644 --- a/src/client/impl/internal/grpc_connections/grpc_connections.h +++ b/src/client/impl/internal/grpc_connections/grpc_connections.h @@ -97,6 +97,13 @@ class TGRpcConnectionsImpl bool TryCreateContext(IQueueClientContextPtr& context); void Stop(bool wait = false); + std::uint64_t GetMaxOutboundMessageSize() const { + if (MaxOutboundMessageSize_ > 0) { + return MaxOutboundMessageSize_; + } + return MaxMessageSize_ > 0 ? MaxMessageSize_ : NGrpc::DEFAULT_GRPC_MESSAGE_SIZE_LIMIT; + } + template using TServiceConnection = NYdbGrpc::TServiceConnection; @@ -136,9 +143,7 @@ class TGRpcConnectionsImpl if (MaxInboundMessageSize_ > 0) { clientConfig.MaxInboundMessageSize = MaxInboundMessageSize_; } - if (MaxOutboundMessageSize_ > 0) { - clientConfig.MaxOutboundMessageSize = MaxOutboundMessageSize_; - } + clientConfig.MaxOutboundMessageSize = GetMaxOutboundMessageSize(); clientConfig.LoadBalancingPolicy = GRpcLoadBalancingPolicy_; diff --git a/src/client/persqueue_public/impl/read_session.cpp b/src/client/persqueue_public/impl/read_session.cpp index 6bbe2df9573..feacfee412f 100644 --- a/src/client/persqueue_public/impl/read_session.cpp +++ b/src/client/persqueue_public/impl/read_session.cpp @@ -167,7 +167,10 @@ void TReadSession::StartClusterDiscovery() { selfShared->OnClusterDiscovery(st, result); }; - auto rpcSettings = TRpcRequestSettings::Make(Settings); + auto rpcSettings = TRpcRequestSettings::Make( + Settings, + {}, + TRpcRequestSettings::TEndpointPolicy::UseDiscoveryEndpoint); rpcSettings.Deadline = TDeadline::AfterDuration(std::chrono::seconds(5)); // TODO: make client timeout setting Connections->RunDeferredScheduleDelayedTask(std::move(cdsRequestCall), TDeadline::SafeDurationCast(delay)); return; @@ -1258,12 +1261,12 @@ void TWriteSessionImpl::UpdateTokenImpl(const NThreading::TFuture& void TWriteSessionImpl::SendImpl() { Y_ABORT_UNLESS(Lock.IsLocked()); - // External cycle splits ready blocks into multiple gRPC messages. Current gRPC message size hard limit is 64MiB + // Split ready blocks into requests bounded by the driver's outbound limit. while(IsReadyToSendNextImpl()) { TClientMessage clientMessage; auto* writeRequest = clientMessage.mutable_write_request(); auto sentAtMs = TInstant::Now().MilliSeconds(); - NGrpc::TRequestSizeLimiter sizeLimiter(2); + NGrpc::TRequestSizeLimiter sizeLimiter(2, NGrpc::GetMaxGrpcMessageSize(*Connections)); // Sent blocks while we can without messages reordering while (IsReadyToSendNextImpl()) { diff --git a/src/client/persqueue_public/ut/read_session_ut.cpp b/src/client/persqueue_public/ut/read_session_ut.cpp index 30e18b65d3c..55fbe94368b 100644 --- a/src/client/persqueue_public/ut/read_session_ut.cpp +++ b/src/client/persqueue_public/ut/read_session_ut.cpp @@ -1906,16 +1906,15 @@ Y_UNIT_TEST_SUITE(ReadSessionImplTest) { false, 0); - std::atomic ready = true; - std::atomic abandoned = false; + std::atomic state = NTopic::EDecompressionTaskState::Ready; - stream->InsertDataEvent(0, 0, data, ready, abandoned); + stream->InsertDataEvent(0, 0, data, state); stream->InsertEvent(TServiceEvent{stream, 0, 0, 0, {}}); - stream->InsertDataEvent(0, 0, data, ready, abandoned); - stream->InsertDataEvent(0, 0, data, ready, abandoned); + stream->InsertDataEvent(0, 0, data, state); + stream->InsertDataEvent(0, 0, data, state); stream->InsertEvent(TServiceEvent{stream, 0, 0, 0, {}}); stream->InsertEvent(TServiceEvent{stream, 0, 0, 0, {}}); - stream->InsertDataEvent(0, 0, data, ready, abandoned); + stream->InsertDataEvent(0, 0, data, state); TDeferredActions actions; diff --git a/src/client/query/client.cpp b/src/client/query/client.cpp index a93d3688875..b76d266e6f6 100644 --- a/src/client/query/client.cpp +++ b/src/client/query/client.cpp @@ -28,9 +28,9 @@ namespace NYdb::inline V3::NQuery { using TQueryObservation = NObservability::TRequestObservation; -NYdb::NRetry::TRetryOperationSettings GetRetrySettings(TDuration timeout, bool isIndempotent) { +NYdb::NRetry::TRetryOperationSettings GetRetrySettings(TDuration timeout, bool isIdempotent) { return NYdb::NRetry::TRetryOperationSettings() - .Idempotent(isIndempotent) + .Idempotent(isIdempotent) .GetSessionClientTimeout(timeout) .MaxTimeout(timeout); } @@ -925,9 +925,9 @@ TStatus TQueryClient::RetryQuerySync(const TQueryWithoutSessionSyncFunc& queryFu } TAsyncExecuteQueryResult TQueryClient::RetryQuery(const std::string& query, const TTxControl& txControl, - TDuration timeout, bool isIndempotent) + TDuration timeout, bool isIdempotent) { - auto settings = GetRetrySettings(timeout, isIndempotent); + auto settings = GetRetrySettings(timeout, isIdempotent); auto queryFunc = [query, txControl](TSession session, TDuration duration) -> TAsyncExecuteQueryResult { return session.ExecuteQuery(query, txControl, TExecuteQuerySettings().ClientTimeout(duration)); }; diff --git a/src/client/table/out.cpp b/src/client/table/out.cpp index 408bfa416cb..3bddb689398 100644 --- a/src/client/table/out.cpp +++ b/src/client/table/out.cpp @@ -54,6 +54,10 @@ Y_DECLARE_OUT_SPEC(, NYdb::NTable::TVectorIndexSettings::EVectorType, stream, va switch (value) { case NYdb::NTable::TVectorIndexSettings::EVectorType::Float: return "float"; + case NYdb::NTable::TVectorIndexSettings::EVectorType::Float16: + return "float16"; + case NYdb::NTable::TVectorIndexSettings::EVectorType::BFloat16: + return "bfloat16"; case NYdb::NTable::TVectorIndexSettings::EVectorType::Uint8: return "uint8"; case NYdb::NTable::TVectorIndexSettings::EVectorType::Int8: diff --git a/src/client/table/table.cpp b/src/client/table/table.cpp index 6d14023a404..cc35a2639f4 100644 --- a/src/client/table/table.cpp +++ b/src/client/table/table.cpp @@ -2898,6 +2898,10 @@ TVectorIndexSettings TVectorIndexSettings::FromProto(const Ydb::Table::VectorInd switch (proto.vector_type()) { case Ydb::Table::VectorIndexSettings::VECTOR_TYPE_FLOAT: return EVectorType::Float; + case Ydb::Table::VectorIndexSettings::VECTOR_TYPE_FLOAT16: + return EVectorType::Float16; + case Ydb::Table::VectorIndexSettings::VECTOR_TYPE_BFLOAT16: + return EVectorType::BFloat16; case Ydb::Table::VectorIndexSettings::VECTOR_TYPE_UINT8: return EVectorType::Uint8; case Ydb::Table::VectorIndexSettings::VECTOR_TYPE_INT8: @@ -2938,6 +2942,10 @@ void TVectorIndexSettings::SerializeTo(Ydb::Table::VectorIndexSettings& settings switch (VectorType) { case EVectorType::Float: return Ydb::Table::VectorIndexSettings::VECTOR_TYPE_FLOAT; + case EVectorType::Float16: + return Ydb::Table::VectorIndexSettings::VECTOR_TYPE_FLOAT16; + case EVectorType::BFloat16: + return Ydb::Table::VectorIndexSettings::VECTOR_TYPE_BFLOAT16; case EVectorType::Uint8: return Ydb::Table::VectorIndexSettings::VECTOR_TYPE_UINT8; case EVectorType::Int8: @@ -3954,7 +3962,9 @@ std::optional TTtlTierSettings::FromProto(const Ydb::Table::Tt action = TTtlDeleteAction(); break; case Ydb::Table::TtlTier::kEvictToExternalStorage: - action = TTtlEvictToExternalStorageAction(tier.evict_to_external_storage().storage()); + action = TTtlEvictToExternalStorageAction(tier.evict_to_external_storage().storage(), + tier.evict_to_external_storage().has_object_key_prefix() + ? std::make_optional(tier.evict_to_external_storage().object_key_prefix()) : std::nullopt); break; case Ydb::Table::TtlTier::ACTION_NOT_SET: return std::nullopt; @@ -4071,8 +4081,23 @@ TTtlEvictToExternalStorageAction::TTtlEvictToExternalStorageAction(const std::st : Storage_(storageName) {} +TTtlEvictToExternalStorageAction::TTtlEvictToExternalStorageAction( + const std::string& storageName, const std::optional& objectKeyPrefix) + : Storage_(storageName) + , ObjectKeyPrefix_(objectKeyPrefix) +{} + void TTtlEvictToExternalStorageAction::SerializeTo(Ydb::Table::EvictionToExternalStorageSettings& proto) const { proto.set_storage(Storage_); + if (ObjectKeyPrefix_) { + proto.set_object_key_prefix(*ObjectKeyPrefix_); + } else { + proto.clear_object_key_prefix(); + } +} + +const std::optional& TTtlEvictToExternalStorageAction::GetObjectKeyPrefix() const { + return ObjectKeyPrefix_; } std::string TTtlEvictToExternalStorageAction::GetStorage() const { diff --git a/src/client/topic/impl/common.h b/src/client/topic/impl/common.h index 80786e9d097..ec20f8ef2cc 100644 --- a/src/client/topic/impl/common.h +++ b/src/client/topic/impl/common.h @@ -14,6 +14,7 @@ #include +#include #include #include #include @@ -38,8 +39,9 @@ namespace NWriteSessionGrpc { inline constexpr std::string_view PARTITION_KEY_META_KEY = "__partition_key"; -inline size_t GetMaxGrpcMessageSize() { - return 120_MB; +inline size_t GetMaxGrpcMessageSize(const TGRpcConnectionsImpl& connections) { + // Keep the existing batching cap, but never exceed the driver's send limit. + return std::min(120_MB, connections.GetMaxOutboundMessageSize()); } using TWireFormatLite = google::protobuf::internal::WireFormatLite; @@ -106,7 +108,7 @@ inline size_t ProtoMessageFieldSize(ui32 fieldNumber, size_t size) { class TRequestSizeLimiter { public: - explicit TRequestSizeLimiter(ui32 envelopeFieldNumber, size_t maxSize = GetMaxGrpcMessageSize()) + explicit TRequestSizeLimiter(ui32 envelopeFieldNumber, size_t maxSize) : EnvelopeFieldNumber(envelopeFieldNumber) , MaxSize(maxSize) { diff --git a/src/client/topic/impl/producer.cpp b/src/client/topic/impl/producer.cpp index 13055bfbcea..b616c87d284 100644 --- a/src/client/topic/impl/producer.cpp +++ b/src/client/topic/impl/producer.cpp @@ -641,6 +641,9 @@ void TProducer::TEventsWorker::SubscribeToPartition(std::uint32_t partition) { } if (partitionIt->second.IsSplitted() || Producer->SplittedPartitionWorkers.contains(partition)) { + std::lock_guard lock(Lock); + SubscribedPartitions.erase(partition); + ReadyFutures.erase(partition); partitionIt->second.Future(NThreading::MakeFuture()); return; } @@ -650,6 +653,13 @@ void TProducer::TEventsWorker::SubscribeToPartition(std::uint32_t partition) { std::weak_ptr producer = Producer->shared_from_this(); std::weak_ptr self = shared_from_this(); + { + // Arm the subscription before Subscribe(): WaitEvent() may run the callback + // synchronously, and that callback takes Lock itself. + std::lock_guard lock(Lock); + SubscribedPartitions.insert(partition); + } + newFuture.Subscribe([self, producer, partition](const NThreading::TFuture&) { auto producerPtr = producer.lock(); if (!producerPtr) { @@ -663,11 +673,25 @@ void TProducer::TEventsWorker::SubscribeToPartition(std::uint32_t partition) { { std::lock_guard lock(selfPtr->Lock); + if (!selfPtr->SubscribedPartitions.contains(partition)) { + return; + } selfPtr->ReadyFutures.insert(partition); } producerPtr->RunMainWorker(static_cast(partition)); }); - partitionIt->second.Future(newFuture); + + { + std::lock_guard lock(Lock); + if (!SubscribedPartitions.contains(partition)) { + return; + } + // The callback above may have run synchronously and mutated Partitions. + auto subscribedPartition = Producer->Partitions.find(partition); + if (subscribedPartition != Producer->Partitions.end()) { + subscribedPartition->second.Future(newFuture); + } + } } std::optional TProducer::TEventsWorker::GetSessionClosedEvent() { @@ -870,6 +894,11 @@ NThreading::TFuture TProducer::TEventsWorker::WaitEvent() { } void TProducer::TEventsWorker::UnsubscribeFromPartition(std::uint32_t partition) { + // SubscribeToPartition's WaitEvent callback inserts into ReadyFutures on the gRPC + // thread. Clear the subscription under the same Lock so that insert cannot race + // with erase, and a callback that already lost the race does not resurrect the partition. + std::lock_guard lock(Lock); + SubscribedPartitions.erase(partition); ReadyFutures.erase(partition); auto partitionIt = Producer->Partitions.find(partition); if (partitionIt != Producer->Partitions.end()) { diff --git a/src/client/topic/impl/producer.h b/src/client/topic/impl/producer.h index 29447426a3d..4ee678baaf0 100644 --- a/src/client/topic/impl/producer.h +++ b/src/client/topic/impl/producer.h @@ -314,6 +314,10 @@ class TProducer : public IProducer, TProducer* Producer; std::unordered_set ReadyFutures; + // Partitions with an armed WaitEvent callback. Both this set and ReadyFutures + // are mutated only under Lock. UnsubscribeFromPartition drops the partition + // from here so an in-flight callback cannot insert into ReadyFutures afterwards. + std::unordered_set SubscribedPartitions; std::unordered_map> PartitionsEventQueues; std::list EventsOutputQueue; diff --git a/src/client/topic/impl/read_session.cpp b/src/client/topic/impl/read_session.cpp index 3c8c536ed16..c145ae05d44 100644 --- a/src/client/topic/impl/read_session.cpp +++ b/src/client/topic/impl/read_session.cpp @@ -208,6 +208,7 @@ bool TReadSession::Close(TDuration timeout) { std::shared_ptr>> dumpCountersContextToCancel; TInstant closeDeadline; bool result = false; + bool zeroTimeout = false; { TDeferredActions deferred; with_lock(Lock) { @@ -217,58 +218,74 @@ bool TReadSession::Close(TDuration timeout) { if (!timeout) { AbortImpl(EStatus::ABORTED, "Close with zero timeout", deferred); - return false; + zeroTimeout = true; + } else { + Closing = true; + session = CbContext->TryGet(); } - - Closing = true; - session = CbContext->TryGet(); } - session->Close(callback); - - callback(); // For the case when there are no subsessions yet. + if (zeroTimeout) { + // deferred posts callbacks as this block ends, then we drop the + // stream -> session cycle without waiting for ~TReadSession. + } else { + session->Close(callback); - auto timeoutCallback = [=](bool) mutable { - promise.TrySetValue(false); - }; + callback(); // For the case when there are no subsessions yet. - auto timeoutContext = Connections->CreateContext(); - if (!timeoutContext) { - AbortImpl(EStatus::ABORTED, DRIVER_IS_STOPPING_DESCRIPTION, deferred); - return false; - } - closeDeadline = TInstant::Now() + timeout; - Connections->ScheduleCallback(timeout, - std::move(timeoutCallback), - timeoutContext); - - // Wait. - NThreading::TFuture resultFuture = promise.GetFuture(); - result = resultFuture.GetValueSync(); - if (result) { - Cancel(timeoutContext); - - NYdb::NIssue::TIssues issues; - issues.AddIssue("Session was gracefully closed"); - EventsQueue->Close(TSessionClosedEvent(EStatus::SUCCESS, std::move(issues)), deferred); - } else { - ++*Settings.Counters_->Errors; - session->Abort(); + auto timeoutCallback = [=](bool) mutable { + promise.TrySetValue(false); + }; - NYdb::NIssue::TIssues issues; - issues.AddIssue(TStringBuilder() << "Session was closed after waiting " << timeout); - EventsQueue->Close(TSessionClosedEvent(EStatus::TIMEOUT, std::move(issues)), deferred); - } - { - std::lock_guard guard(Lock); - Aborting = true; // Set abort flag for doing nothing on destructor. - cbContextToCancel = CbContext; - dumpCountersContextToCancel = DumpCountersContext; + auto timeoutContext = Connections->CreateContext(); + if (!timeoutContext) { + AbortImpl(EStatus::ABORTED, DRIVER_IS_STOPPING_DESCRIPTION, deferred); + return false; + } + closeDeadline = TInstant::Now() + timeout; + Connections->ScheduleCallback(timeout, + std::move(timeoutCallback), + timeoutContext); + + // Wait. + NThreading::TFuture resultFuture = promise.GetFuture(); + result = resultFuture.GetValueSync(); + if (result) { + Cancel(timeoutContext); + + NYdb::NIssue::TIssues issues; + issues.AddIssue("Session was gracefully closed"); + EventsQueue->Close(TSessionClosedEvent(EStatus::SUCCESS, std::move(issues)), deferred); + } else { + ++*Settings.Counters_->Errors; + session->Abort(); + + NYdb::NIssue::TIssues issues; + issues.AddIssue(TStringBuilder() << "Session was closed after waiting " << timeout); + EventsQueue->Close(TSessionClosedEvent(EStatus::TIMEOUT, std::move(issues)), deferred); + } + { + std::lock_guard guard(Lock); + Aborting = true; // Set abort flag for doing nothing on destructor. + cbContextToCancel = CbContext; + dumpCountersContextToCancel = DumpCountersContext; + } + if (!session->WaitAllDecompressionTasks(closeDeadline)) { + LOG_LAZY(Log, TLOG_WARNING, GetLogPrefix() << "Some decompression tasks are still running after read session close timeout"); + } + ClearAllEvents(); + session->ClearAllPartitionStreamEvents(); } - if (!session->WaitAllDecompressionTasks(closeDeadline)) { - LOG_LAZY(Log, TLOG_WARNING, GetLogPrefix() << "Some decompression tasks are still running after read session close timeout"); + } + if (zeroTimeout) { + if (CbContext) { + if (auto abortedSession = CbContext->LockShared()) { + ClearAllEvents(); + abortedSession->ClearAllPartitionStreamEvents(); + } else { + ClearAllEvents(); + } } - ClearAllEvents(); - session->ClearAllPartitionStreamEvents(); + return false; } if (cbContextToCancel) { cbContextToCancel->Cancel(); diff --git a/src/client/topic/impl/read_session_impl.h b/src/client/topic/impl/read_session_impl.h index ee994937b4f..21bd1279392 100644 --- a/src/client/topic/impl/read_session_impl.h +++ b/src/client/topic/impl/read_session_impl.h @@ -116,6 +116,13 @@ class TReadSessionEventsQueue; class TReadSession; +enum class EDecompressionTaskState : ui8 { + InProcess, + Cleanup, + Ready, + Abandoned, +}; + template class TDataDecompressionInfo; @@ -308,7 +315,8 @@ class TDataDecompressionInfo : public std::enable_shared_from_this ret; for (auto i = ReadyThresholds.begin(), end = ReadyThresholds.end(); i != end; ++i) { - if (i->Ready) { + const auto state = i->State.load(); + if (state == EDecompressionTaskState::Ready || state == EDecompressionTaskState::Abandoned) { ret.first = i->Batch; ret.second = i->Message; ++readyCount; @@ -354,8 +362,7 @@ class TDataDecompressionInfo : public std::enable_shared_from_this Ready = false; - std::atomic Abandoned = false; // Marked by true either when decompression is completed or message is not needed anymore + std::atomic State = EDecompressionTaskState::InProcess; }; struct TDecompressionTask { @@ -433,26 +440,30 @@ class TDataDecompressionInfo : public std::enable_shared_from_this class TDataDecompressionEvent { public: - TDataDecompressionEvent(size_t batch, size_t message, TDataDecompressionInfoPtr parent, std::atomic& ready, std::atomic& abandoned) : + TDataDecompressionEvent(size_t batch, size_t message, TDataDecompressionInfoPtr parent, std::atomic& state) : Batch{batch}, Message{message}, Parent{std::move(parent)}, - Ready{ready}, - Abandoned{abandoned} + State{state} { } bool IsReady() const { - return Ready; + const auto state = State.load(); + return state == EDecompressionTaskState::Ready || state == EDecompressionTaskState::Abandoned; + } + + bool IsAbandoned() const { + const auto state = State.load(); + return state == EDecompressionTaskState::Cleanup || state == EDecompressionTaskState::Abandoned; } bool SetAbandoned() { - if (bool expected = false; Abandoned.compare_exchange_strong(expected, true)) { + auto expected = EDecompressionTaskState::InProcess; + if (State.compare_exchange_strong(expected, EDecompressionTaskState::Cleanup)) { return true; } - - // If Ready=false, decompression task already successfully cancelled - return !Ready; + return expected != EDecompressionTaskState::Ready; } typename TDataDecompressionInfo::TDecompressedData @@ -470,8 +481,7 @@ class TDataDecompressionEvent { size_t Batch; size_t Message; TDataDecompressionInfoPtr Parent; - std::atomic& Ready; - std::atomic& Abandoned; + std::atomic& State; }; template @@ -536,14 +546,12 @@ struct TRawPartitionStreamEvent { TRawPartitionStreamEvent(size_t batch, size_t message, TDataDecompressionInfoPtr parent, - std::atomic &ready, - std::atomic& abandoned) + std::atomic& state) : Event(std::in_place_type_t>(), batch, message, std::move(parent), - ready, - abandoned) + state) { } @@ -584,6 +592,10 @@ struct TRawPartitionStreamEvent { return std::get>(Event).IsReady(); } + + bool IsAbandoned() const { + return IsDataEvent() && GetDataEvent().IsAbandoned(); + } }; template @@ -634,6 +646,13 @@ class TRawPartitionStreamEventQueue { Ready.clear(); } + // Drop the session ref. An empty queue must not keep + // TSingleClusterReadSessionImpl alive through TCallbackContext. + void ReleaseContext() noexcept { + clear(); + CbContext.reset(); + } + void SignalReadyEvents(TIntrusivePtr> stream, TReadSessionEventsQueue& queue, TDeferredActions& deferred); @@ -789,10 +808,9 @@ class TPartitionStreamImpl : public TAPartitionStream { void InsertDataEvent(size_t batch, size_t message, TDataDecompressionInfoPtr parent, - std::atomic& ready, - std::atomic& abandoned) + std::atomic& state) { - EventsQueue.emplace_back(batch, message, std::move(parent), ready, abandoned); + EventsQueue.emplace_back(batch, message, std::move(parent), state); } bool HasEvents() const { @@ -808,7 +826,7 @@ class TPartitionStreamImpl : public TAPartitionStream { } TCallbackContextPtr GetCbContext() const { - return CbContext; + return CopyCallbackContext(); } TLog GetLog() const; @@ -880,7 +898,18 @@ class TPartitionStreamImpl : public TAPartitionStream { } TRawPartitionStreamEventQueue ExtractQueue() noexcept { - return std::exchange(EventsQueue, TRawPartitionStreamEventQueue(CbContext)); + return std::exchange(EventsQueue, TRawPartitionStreamEventQueue(CopyCallbackContext())); + } + + // Breaks stream -> callback context -> session -> stream. Called when the + // session is closing; later callbacks must not use this stream. + void DropCallbackContext() noexcept { + EventsQueue.ReleaseContext(); + std::atomic_store(&CbContext, TCallbackContextPtr{}); + } + + TCallbackContextPtr CopyCallbackContext() const { + return std::atomic_load(&CbContext); } static void GetDataEventImpl(TIntrusivePtr> partitionStream, @@ -1004,8 +1033,7 @@ class TReadSessionEventsQueue: public TBaseSessionEventsQueue parent, - std::atomic& ready, - std::atomic& abandoned); + std::atomic& state); void SignalEventImpl(TIntrusivePtr> partitionStream, TDeferredActions& deferred, @@ -1023,6 +1051,22 @@ class TReadSessionEventsQueue: public TBaseSessionEventsQueueGetLock(). Then this takes Mutex: same order as + // SignalReadyEvents. PushEvent and GetEvent mutate the stream queue under + // Mutex alone, so ExtractQueue/DropCallbackContext must take it too. + void ExtractPartitionStreamQueue( + const TIntrusivePtr>& stream, + std::vector>& deferredDelete) + { + std::lock_guard guard(TParent::Mutex); + if (stream->HasEvents()) { + deferredDelete.push_back(stream->ExtractQueue()); + } + // ExtractQueue installs a fresh queue that still owns the session. + // Drop it too, or the stream keeps the session (and itself) alive. + stream->DropCallbackContext(); + } + void SetCallbackContext(TCallbackContextPtr& ctx) { CbContext = ctx; } diff --git a/src/client/topic/impl/read_session_impl.ipp b/src/client/topic/impl/read_session_impl.ipp index 3092a5958da..5458fe4ba75 100644 --- a/src/client/topic/impl/read_session_impl.ipp +++ b/src/client/topic/impl/read_session_impl.ipp @@ -64,8 +64,11 @@ static const bool DecompressEverything = !std::string{std::getenv("PQ_DECOMPRESS template TLog TPartitionStreamImpl::GetLog() const { - if (auto session = CbContext->LockShared()) { - return session->GetLog(); + auto callbackContext = CopyCallbackContext(); + if (callbackContext) { + if (auto session = callbackContext->LockShared()) { + return session->GetLog(); + } } return {}; } @@ -73,7 +76,11 @@ TLog TPartitionStreamImpl::GetLog() const { template void TPartitionStreamImpl::Commit(uint64_t startOffset, uint64_t endOffset) { std::vector> toCommit; - if (auto sessionShared = CbContext->LockShared()) { + auto callbackContext = CopyCallbackContext(); + if (!callbackContext) { + return; + } + if (auto sessionShared = callbackContext->LockShared()) { Y_ABORT_UNLESS(endOffset > startOffset); { std::lock_guard guard(sessionShared->Lock); @@ -95,14 +102,22 @@ void TPartitionStreamImpl::Commit(uint64_t startOffset, ui template void TPartitionStreamImpl::RequestStatus() { - if (auto sessionShared = CbContext->LockShared()) { + auto callbackContext = CopyCallbackContext(); + if (!callbackContext) { + return; + } + if (auto sessionShared = callbackContext->LockShared()) { sessionShared->RequestPartitionStreamStatus(this); } } template void TPartitionStreamImpl::ConfirmCreate(std::optional readOffset, std::optional commitOffset, std::optional maxOffset) { - if (auto sessionShared = CbContext->LockShared()) { + auto callbackContext = CopyCallbackContext(); + if (!callbackContext) { + return; + } + if (auto sessionShared = callbackContext->LockShared()) { if (commitOffset.has_value()) { SetFirstNotReadOffset(commitOffset.value()); } @@ -112,14 +127,22 @@ void TPartitionStreamImpl::ConfirmCreate(std::optional void TPartitionStreamImpl::ConfirmDestroy() { - if (auto sessionShared = CbContext->LockShared()) { + auto callbackContext = CopyCallbackContext(); + if (!callbackContext) { + return; + } + if (auto sessionShared = callbackContext->LockShared()) { sessionShared->ConfirmPartitionStreamDestroy(this); } } template void TPartitionStreamImpl::ConfirmEnd(std::span childIds) { - if (auto sessionShared = CbContext->LockShared()) { + auto callbackContext = CopyCallbackContext(); + if (!callbackContext) { + return; + } + if (auto sessionShared = callbackContext->LockShared()) { sessionShared->ConfirmPartitionStreamEnd(this, childIds); } } @@ -175,6 +198,9 @@ void TRawPartitionStreamEventQueue::SignalReadyEvents(TInt TDeferredActions& deferred) { if constexpr (!UseMigrationProtocol) { + if (!CbContext) { + return; + } if (auto session = CbContext->LockShared()) { if (!session->AllParentSessionsHasBeenRead(stream->GetPartitionId(), stream->GetPartitionSessionId())) { return; @@ -189,9 +215,14 @@ void TRawPartitionStreamEventQueue::SignalReadyEvents(TInt NotReady.pop_front(); }; - while (!NotReady.empty() && NotReady.front().IsReady()) { + while (!NotReady.empty() && (NotReady.front().IsReady() || NotReady.front().IsAbandoned())) { auto& front = NotReady.front(); + if (front.IsAbandoned()) { + NotReady.pop_front(); + continue; + } + if (front.IsDataEvent()) { if (queue.HasDataEventCallback()) { std::vector::TMessage> messages; @@ -239,9 +270,7 @@ void TRawPartitionStreamEventQueue::DeleteNotReadyTail(TDe for (auto& event : NotReady) { const bool isDataEvent = event.IsDataEvent(); - if (event.IsReady() || - (isDataEvent && !event.GetDataEvent().SetAbandoned()) // Try to cancel inflight decompression tasks if any (returns true if message was decompressed and become ready) - ) { + if (!isDataEvent || !event.GetDataEvent().SetAbandoned()) { if (!hasNonReadyEvents) { // Continue ready events prefix ready.push_back(std::move(event)); @@ -280,7 +309,7 @@ void TRawPartitionStreamEventQueue::Cleanup(TDeferredActio } auto& dataEvent = event.GetDataEvent(); - if (event.IsReady() || !dataEvent.SetAbandoned()) { + if (!dataEvent.SetAbandoned()) { accumulator.Add(dataEvent.GetParent(), dataEvent.GetDataSize(), dataEvent.GetMessageCount()); } else { infos.push_back(dataEvent.GetParent()); @@ -1959,9 +1988,7 @@ void TSingleClusterReadSessionImpl::ClearAllPartitionStrea deferredDelete.reserve(streams.size()); for (auto& stream : streams) { std::lock_guard guard(stream->GetLock()); - if (stream->HasEvents()) { - deferredDelete.push_back(stream->ExtractQueue()); - } + EventsQueue->ExtractPartitionStreamQueue(stream, deferredDelete); } for (auto& queue : deferredDelete) { @@ -2677,15 +2704,14 @@ bool TReadSessionEventsQueue::PushDataEvent(TIntrusivePtr< size_t batch, size_t message, TDataDecompressionInfoPtr parent, - std::atomic& ready, - std::atomic& abandoned) + std::atomic& state) { std::lock_guard guard(TParent::Mutex); if (this->Closed) { return false; } - partitionStream->InsertDataEvent(batch, message, parent, ready, abandoned); + partitionStream->InsertDataEvent(batch, message, parent, state); return true; } @@ -2722,7 +2748,7 @@ void TRawPartitionStreamEventQueue::GetDataEventImpl(TIntr auto& front = queue.front(); - return front.IsDataEvent() && front.IsReady(); + return front.IsDataEvent() && front.IsReady() && !front.IsAbandoned(); }; Y_ABORT_UNLESS(readyDataInTheHead()); @@ -3203,8 +3229,7 @@ bool TDataDecompressionInfo::PlanDecompressionTasks(double CurrentDecompressingMessage.first, CurrentDecompressingMessage.second, TDataDecompressionInfo::shared_from_this(), - ReadyThresholds.back().Ready, - ReadyThresholds.back().Abandoned); + ReadyThresholds.back().State); if (!pushRes) { deferred.DeferDestroyDecompressionInfos({TDataDecompressionInfo::shared_from_this()}); session->AbortImpl(&deferred); @@ -3290,7 +3315,7 @@ TDataDecompressionInfo::BuildDecompressedData(TIntrusivePt continue; } seqNo = static_cast(*codecResult.BatchBaseSequence) + static_cast(recordMeta.SequenceDelta); - createTime = TInstant::MilliSeconds(*codecResult.BatchBaseTimestampMs + recordMeta.TimestampDelta); + createTime = TInstant::MilliSeconds(NKafka::GetRecordTimestamp(*codecResult.BatchBaseTimestampMs, recordMeta.TimestampDelta)); } TReadSessionEvent::TDataReceivedEvent::TMessageInformation messageInfo( @@ -3641,15 +3666,16 @@ void TDataDecompressionInfo::TDecompressionTask::operator( parent->OnDataDecompressed(SourceDataSize, EstimatedDecompressedSize, DecompressedSize, messagesProcessed); parent->SourceDataNotProcessed -= dataProcessed; - Ready->Ready = true; - - if (auto session = parent->CbContext->LockShared()) { - session->GetEventsQueue()->SignalReadyEvents(PartitionStream); - } - - if (bool expected = false; !Ready->Abandoned.compare_exchange_strong(expected, true)) { - // Message is dropped due to partition stream cancellation, we should release decompressed memory + auto expected = EDecompressionTaskState::InProcess; + if (Ready->State.compare_exchange_strong(expected, EDecompressionTaskState::Ready)) { + if (auto session = parent->CbContext->LockShared()) { + session->GetEventsQueue()->SignalReadyEvents(PartitionStream); + } + } else { + Y_ABORT_UNLESS(expected == EDecompressionTaskState::Cleanup); + // Cleanup claimed the whole task before it became ready. parent->OnUserRetrievedEvent(DecompressedSize, messagesProcessed); + Ready->State.store(EDecompressionTaskState::Abandoned); } if (auto session = parent->CbContext->LockShared()) { diff --git a/src/client/topic/impl/topic.cpp b/src/client/topic/impl/topic.cpp index 590a3f92f7b..58de1118632 100644 --- a/src/client/topic/impl/topic.cpp +++ b/src/client/topic/impl/topic.cpp @@ -634,6 +634,11 @@ TAsyncStatus TTopicClient::CommitOffset(const std::string& path, uint64_t partit return Impl_->CommitOffset(path, partitionId, consumerName, offset, settings); } +TAsyncStatus TTopicClient::ResetOffset(const std::string& path, const std::string& consumerName, + const TResetOffsetSettings& settings) { + return Impl_->ResetOffset(path, consumerName, settings); +} + namespace { Ydb::Topic::SupportedCodecs SerializeCodecs(const std::vector& codecs) { diff --git a/src/client/topic/impl/topic_impl.h b/src/client/topic/impl/topic_impl.h index 406b9c5901c..44e5740b005 100644 --- a/src/client/topic/impl/topic_impl.h +++ b/src/client/topic/impl/topic_impl.h @@ -13,6 +13,7 @@ #include #include +#include namespace NYdb::inline V3::NTopic { struct TOffsetsRange { @@ -288,6 +289,32 @@ class TTopicClient::TImpl : public TClientImplCommon { TRpcRequestSettings::Make(settings)); } + TAsyncStatus ResetOffset(const std::string& path, const std::string& consumerName, + const TResetOffsetSettings& settings) + { + Ydb::Topic::ResetOffsetRequest request = MakeOperationRequest(settings); + request.set_path(TStringType{path}); + request.set_consumer(TStringType{consumerName}); + switch (settings.Position_) { + case TResetOffsetSettings::EPosition::Earliest: + request.mutable_earliest(); + break; + case TResetOffsetSettings::EPosition::Latest: + request.mutable_latest(); + break; + case TResetOffsetSettings::EPosition::FromWrittenAt: + *request.mutable_from_written_at()->mutable_written_at() = + ::google::protobuf::util::TimeUtil::MillisecondsToTimestamp(settings.FromWrittenAt_.MilliSeconds()); + break; + case TResetOffsetSettings::EPosition::Unspecified: + break; + } + return RunSimple( + std::move(request), + &Ydb::Topic::V1::TopicService::Stub::AsyncResetOffset, + TRpcRequestSettings::Make(settings)); + } + TAsyncStatus UpdateOffsetsInTransaction(const TTransactionId& tx, const std::vector& topics, const std::string& consumerName, diff --git a/src/client/topic/impl/write_session_impl.cpp b/src/client/topic/impl/write_session_impl.cpp index 1e230cb5f9c..6631300cfb1 100644 --- a/src/client/topic/impl/write_session_impl.cpp +++ b/src/client/topic/impl/write_session_impl.cpp @@ -1990,14 +1990,14 @@ void TWriteSessionImpl::SendStandardBlock( void TWriteSessionImpl::SendImpl() { Y_ABORT_UNLESS(Lock.IsLocked()); - // External cycle splits ready blocks into multiple gRPC messages. Current gRPC message size hard limit is 64MiB. + // Split ready blocks into requests bounded by the driver's outbound limit. while (IsReadyToSendNextImpl()) { TClientMessage clientMessage; auto* writeRequest = clientMessage.mutable_write_request(); ui32 prevCodec = 0; - NGrpc::TRequestSizeLimiter sizeLimiter(2); + NGrpc::TRequestSizeLimiter sizeLimiter(2, NGrpc::GetMaxGrpcMessageSize(*Connections)); // Send blocks while we can without messages reordering. while (IsReadyToSendNextImpl()) { diff --git a/src/client/topic/ut/read_session_accounting_ut.cpp b/src/client/topic/ut/read_session_accounting_ut.cpp new file mode 100644 index 00000000000..1a49aa23c38 --- /dev/null +++ b/src/client/topic/ut/read_session_accounting_ut.cpp @@ -0,0 +1,65 @@ +#include +#include + +#include + +namespace NYdb::inline V3::NTopic { + + Y_UNIT_TEST_SUITE(TReadSessionDecompressionAccounting) { + Y_UNIT_TEST(CleanupAndDecompressionReleaseSameMessage) { + TReadSessionSettings settings; + auto counters = MakeIntrusive(); + MakeCountersNotNull(*counters); + settings.MaxMemoryUsageBytes(1_MB).Counters(counters); + + auto events = std::make_shared>(settings); + auto context = MakeWithCallbackContext>( + settings, "", "read-session", "", TLog{}, nullptr, events, nullptr, 1, 1); + auto session = context->TryGet(); + auto partition = MakeIntrusive>( + ui64{1}, "topic", "read-session", 0, 1, 0, std::nullopt, context); + + TPartitionData data; + auto* batch = data.add_batches(); + batch->set_codec(Ydb::Topic::CODEC_RAW); + batch->add_message_data()->set_data(std::string(10, 'a')); + batch->add_message_data()->set_data(std::string(20, 'b')); + auto info = std::make_shared>( + std::move(data), context, true); + + { + TDeferredActions actions; + UNIT_ASSERT(info->PlanDecompressionTasks(1.0, partition, actions)); + } + + auto queue = partition->ExtractQueue(); + UNIT_ASSERT_VALUES_EQUAL(queue.size(), 2); + TRawPartitionStreamEvent first(std::move(queue.front())); + queue.pop_front(); + TRawPartitionStreamEvent second(std::move(queue.front())); + queue.pop_front(); + queue.emplace_back(std::move(first)); + + { + TDeferredActions actions; + const auto estimatedSize = info->StartDecompressionTasks( + std::make_shared(), 1_MB, actions); + UNIT_ASSERT_VALUES_EQUAL(estimatedSize, 30); + // StartDecompressionTasksImpl normally reserves this estimated size. + session->OnDataDecompressed(0, 0, estimatedSize, 0); + + // Cleanup marks the first message abandoned while the task is pending. + queue.Cleanup(actions); + } // The task releases all 30 bytes claimed by cleanup. + + UNIT_ASSERT(second.IsReady()); + queue.emplace_back(std::move(second)); + { + TDeferredActions actions; + // Cleanup must not release this message again after the worker released the task. + queue.Cleanup(actions); + } + } + } // Y_UNIT_TEST_SUITE(TReadSessionDecompressionAccounting) + +} // namespace NYdb::inline V3::NTopic diff --git a/src/client/topic/ut/read_session_kafka_timestamps_ut.cpp b/src/client/topic/ut/read_session_kafka_timestamps_ut.cpp new file mode 100644 index 00000000000..681a74f1c5e --- /dev/null +++ b/src/client/topic/ut/read_session_kafka_timestamps_ut.cpp @@ -0,0 +1,96 @@ +#include "ut_utils/topic_sdk_test_setup.h" + +#include +#include +#include + +#include + +#include + +#include +#include +#include +#include +#include + +namespace NYdb::inline V3::NTopic::NTests { + + Y_UNIT_TEST_SUITE(ReadSessionKafkaTimestamps) { + Y_UNIT_TEST(ReadKafkaBatchesWithWrappingTimestamps) { + TTopicSdkTestSetup setup{TEST_CASE_NAME}; + auto client = setup.MakeClient(); + const TDuration timeout = TDuration::Seconds(30); + // Store Kafka bytes under a test codec so the server does not cut the + // batch before it reaches the reader's Kafka metadata handling. + TCodecMap::GetTheCodecMap().Set(static_cast(ECodec::CUSTOM), std::make_unique()); + const auto altered = client.AlterTopic(setup.GetTopicPath(), TAlterTopicSettings() + .SetSupportedCodecs({ECodec::CUSTOM}) + .ClientTimeout(timeout)) + .GetValueSync(); + UNIT_ASSERT_VALUES_EQUAL_C(altered.GetStatus(), EStatus::SUCCESS, altered.GetIssues().ToString()); + + constexpr i64 minTimestamp = std::numeric_limits::min(); + constexpr i64 maxTimestamp = std::numeric_limits::max(); + const std::vector payloads = {"small message", std::string(512_KB + 1, 'x')}; + std::vector expectedTimestamps; + auto writeSession = client.CreateSimpleBlockingWriteSession(TWriteSessionSettings() + .Path(setup.GetTopicPath()) + .ProducerId("timestamp-producer") + .MessageGroupId("timestamp-producer") + .Codec(ECodec::RAW)); + i32 nextSequence = 1; + for (const auto compression : {NKafka::ECompressionType::NONE, NKafka::ECompressionType::GZIP, NKafka::ECompressionType::ZSTD}) { + for (const i64 baseTimestamp : {minTimestamp, maxTimestamp}) { + for (size_t i = 0; i < payloads.size(); ++i) { + NKafka::TKafkaRecordBatch batch; + batch.Magic = 2; + batch.Attributes = static_cast(compression); + batch.ProducerId = 42; + batch.ProducerEpoch = 0; + batch.BaseSequence = nextSequence; + batch.BaseTimestamp = baseTimestamp; + batch.MaxTimestamp = maxTimestamp; + NKafka::TKafkaRecord record; + record.OffsetDelta = 0; + record.TimestampDelta = i == 0 ? 0 : (baseTimestamp == minTimestamp ? -1 : 1); + record.SetValue(TString(payloads[i])); + record.Length = record.Size(2) - NKafka::NPrivate::SizeOfVarint(0); + batch.Records.push_back(std::move(record)); + + // Expected Java long addition, independent of GetRecordTimestamp. + const i64 timestamp = i == 0 ? baseTimestamp : (baseTimestamp == minTimestamp ? maxTimestamp : minTimestamp); + expectedTimestamps.push_back(TInstant::MilliSeconds(static_cast(timestamp))); + const TString bytes = NKafka::WriteKafkaRecordBatch(batch); + auto message = TWriteMessage::CompressedMessage( + std::string_view(bytes.data(), bytes.size()), ECodec::CUSTOM, payloads[i].size()); + message.SeqNo(nextSequence); + UNIT_ASSERT(writeSession->Write(std::move(message), nullptr, timeout)); + ++nextSequence; + } + } + } + UNIT_ASSERT(writeSession->Close(timeout)); + + size_t received = 0; + auto result = setup.Read(setup.GetTopicPath(), setup.GetConsumerName(), + [&](TReadSessionEvent::TDataReceivedEvent& event) { + for (const auto& message : event.GetMessages()) { + UNIT_ASSERT(received < expectedTimestamps.size()); + UNIT_ASSERT(!message.HasException()); + UNIT_ASSERT_VALUES_EQUAL(message.GetData(), payloads[received % payloads.size()]); + UNIT_ASSERT_VALUES_EQUAL(message.GetOffset(), received); + UNIT_ASSERT_VALUES_EQUAL(message.GetSeqNo(), received + 1); + UNIT_ASSERT_VALUES_EQUAL(message.GetCreateTime(), expectedTimestamps[received]); + ++received; + } + return received < expectedTimestamps.size(); + }, std::nullopt, timeout); + UNIT_ASSERT(!result.Timeout); + UNIT_ASSERT_VALUES_EQUAL(received, expectedTimestamps.size()); + UNIT_ASSERT(result.Reader->Close(timeout)); + UNIT_ASSERT_VALUES_EQUAL(result.Reader->GetCounters()->MessagesRead->Val(), received); + } + } // Y_UNIT_TEST_SUITE(ReadSessionKafkaTimestamps) + +} // namespace NYdb::inline V3::NTopic::NTests diff --git a/src/client/topic/ut/reset_offset_ut.cpp b/src/client/topic/ut/reset_offset_ut.cpp new file mode 100644 index 00000000000..b0ba2110e8c --- /dev/null +++ b/src/client/topic/ut/reset_offset_ut.cpp @@ -0,0 +1,233 @@ +#include "ut_utils/topic_sdk_test_setup.h" + +#include +#include + +#include + +using namespace NYdb; +using namespace NYdb::NTopic; +using namespace NYdb::NTopic::NTests; +using namespace NKikimr::NPQ::NTest; + +namespace { + +ui64 GetCommittedOffset(TTopicSdkTestSetup& setup, const TString& topic, const TString& consumer, ui32 partitionId = 0) { + auto describe = setup.DescribeConsumer(topic, consumer); + UNIT_ASSERT_LT(partitionId, describe.GetPartitions().size()); + const auto& stats = describe.GetPartitions()[partitionId].GetPartitionConsumerStats(); + UNIT_ASSERT(stats); + return stats->GetCommittedOffset(); +} + +} // namespace + +Y_UNIT_TEST_SUITE(TResetOffsetSdkTests) { + +Y_UNIT_TEST(EarliestLatestTimestamp) { + TTopicSdkTestSetup setup("ResetOffsetSdk", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + setup.Write("topic1", "m1", 0); + setup.Write("topic1", "m2", 0); + + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath("topic1"); + + { + auto status = client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 2); + } + { + auto status = client.ResetOffset(path, "consumer", TResetOffsetSettings().Earliest()).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 0); + } + { + auto status = client.ResetOffset(path, "consumer", TResetOffsetSettings().FromWrittenAt(TInstant::Now() + TDuration::Hours(1))).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 2); + } +} + +Y_UNIT_TEST(MissingTopicAndConsumer) { + TTopicSdkTestSetup setup("ResetOffsetSdkErrors", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + TTopicClient client(setup.MakeDriver()); + + { + auto status = client.ResetOffset("/Root/missing", "consumer", TResetOffsetSettings().Earliest()).GetValueSync(); + UNIT_ASSERT(!status.IsSuccess()); + UNIT_ASSERT_VALUES_EQUAL(status.GetStatus(), EStatus::SCHEME_ERROR); + } + { + auto status = client.ResetOffset(setup.GetFullTopicPath("topic1"), "no-such-consumer", TResetOffsetSettings().Earliest()).GetValueSync(); + UNIT_ASSERT(!status.IsSuccess()); + UNIT_ASSERT_VALUES_EQUAL(status.GetStatus(), EStatus::SCHEME_ERROR); + } +} + +Y_UNIT_TEST(UnspecifiedPositionRejected) { + TTopicSdkTestSetup setup("ResetOffsetSdkNoPosition", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + TTopicClient client(setup.MakeDriver()); + auto status = client.ResetOffset(setup.GetFullTopicPath("topic1"), "consumer", TResetOffsetSettings()).GetValueSync(); + UNIT_ASSERT(!status.IsSuccess()); + UNIT_ASSERT_VALUES_EQUAL(status.GetStatus(), EStatus::BAD_REQUEST); +} + +Y_UNIT_TEST(IdempotentLatest) { + TTopicSdkTestSetup setup("ResetOffsetSdkIdempotent", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + setup.Write("topic1", "m1", 0); + setup.Write("topic1", "m2", 0); + + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath("topic1"); + UNIT_ASSERT_C(client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync().IsSuccess(), + "first latest"); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 2); + UNIT_ASSERT_C(client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync().IsSuccess(), + "second latest"); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 2); +} + +Y_UNIT_TEST(OtherConsumerUnaffected) { + TTopicSdkTestSetup setup("ResetOffsetSdkTwoConsumers", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer-a"); + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath("topic1"); + auto alter = client.AlterTopic(path, TAlterTopicSettings() + .BeginAddConsumer("consumer-b") + .EndAddConsumer()).GetValueSync(); + UNIT_ASSERT_C(alter.IsSuccess(), alter.GetIssues().ToString()); + + setup.Write("topic1", "m1", 0); + setup.Write("topic1", "m2", 0); + + UNIT_ASSERT_C(client.ResetOffset(path, "consumer-a", TResetOffsetSettings().Latest()).GetValueSync().IsSuccess(), + "reset a"); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer-a"), 2); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer-b"), 0); +} + +Y_UNIT_TEST(ResetLatestAllPartitionsCommitted) { + constexpr ui32 partitionCount = 1024; + TTopicSdkTestSetup setup("ResetOffsetSdkAllPartitions", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer", partitionCount); + + auto client = setup.MakeClient(); + const auto path = setup.GetFullTopicPath("topic1"); + for (ui32 partitionId = 0; partitionId < partitionCount; ++partitionId) { + auto session = client.CreateSimpleBlockingWriteSession( + TWriteSessionSettings() + .Path(path) + .PartitionId(partitionId) + .DeduplicationEnabled(false) + .Codec(ECodec::RAW)); + UNIT_ASSERT_C(session->Write("m"), TStringBuilder() << "write partition " << partitionId); + UNIT_ASSERT_C(session->Close(), TStringBuilder() << "close partition " << partitionId); + } + + auto status = client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + + auto descr = setup.DescribeConsumer("topic1", "consumer"); + UNIT_ASSERT_VALUES_EQUAL(descr.GetPartitions().size(), partitionCount); + for (const auto& part : descr.GetPartitions()) { + const auto& stats = part.GetPartitionConsumerStats(); + UNIT_ASSERT_C(stats, TStringBuilder() << "partition " << part.GetPartitionId()); + UNIT_ASSERT_VALUES_EQUAL_C(stats->GetCommittedOffset(), 1, TStringBuilder() << "partition " << part.GetPartitionId()); + } +} + +Y_UNIT_TEST(LatestSurvivesTabletReboot) { + TTopicSdkTestSetup setup("ResetOffsetSdkRebootLatest", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + setup.Write("topic1", "m1", 0); + + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath("topic1"); + auto status = client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + + auto assertCommittedAtEnd = [&] { + auto descr = setup.DescribeConsumer("topic1", "consumer"); + UNIT_ASSERT_VALUES_EQUAL(descr.GetPartitions().size(), 1); + const auto& part = descr.GetPartitions()[0]; + UNIT_ASSERT(part.GetPartitionStats()); + UNIT_ASSERT(part.GetPartitionConsumerStats()); + UNIT_ASSERT_VALUES_EQUAL(part.GetPartitionStats()->GetEndOffset(), 1); + UNIT_ASSERT_VALUES_EQUAL(part.GetPartitionConsumerStats()->GetCommittedOffset(), + part.GetPartitionStats()->GetEndOffset()); + }; + + assertCommittedAtEnd(); + setup.GetServer().KillTopicPqTablets(TString{path}); + assertCommittedAtEnd(); +} + +Y_UNIT_TEST(RewindSurvivesTabletReboot) { + TTopicSdkTestSetup setup("ResetOffsetSdkRebootRewind", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopic("topic1", "consumer"); + setup.Write("topic1", "m1", 0); + setup.Write("topic1", "m2", 0); + + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath("topic1"); + UNIT_ASSERT_C(client.ResetOffset(path, "consumer", TResetOffsetSettings().Latest()).GetValueSync().IsSuccess(), + "latest"); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 2); + UNIT_ASSERT_C(client.ResetOffset(path, "consumer", TResetOffsetSettings().Earliest()).GetValueSync().IsSuccess(), + "earliest"); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 0); + + setup.GetServer().KillTopicPqTablets(TString{path}); + UNIT_ASSERT_VALUES_EQUAL(GetCommittedOffset(setup, "topic1", "consumer"), 0); +} + +Y_UNIT_TEST(RewindInactiveAfterSplit) { + TTopicSdkTestSetup setup("ResetOffsetSdkInactive", TTopicSdkTestSetup::MakeServerSettings(), false); + setup.CreateTopicWithAutoscale(TEST_TOPIC, TEST_CONSUMER, 1, 100); + setup.Write(TEST_TOPIC, "before-split", 0); + + TTopicClient client(setup.MakeDriver()); + const auto path = setup.GetFullTopicPath(TEST_TOPIC); + auto committed = client.ResetOffset(path, TEST_CONSUMER, TResetOffsetSettings().Latest()).GetValueSync(); + UNIT_ASSERT_C(committed.IsSuccess(), committed.GetIssues().ToString()); + + ui64 txId = 1000; + SplitPartition(setup, txId, 0, "\x80"); + + for (int i = 0; i < 50; ++i) { + auto descr = setup.DescribeConsumer(TEST_TOPIC, TEST_CONSUMER); + bool hasInactive = false; + for (const auto& part : descr.GetPartitions()) { + if (!part.GetActive()) { + hasInactive = true; + break; + } + } + if (hasInactive) { + break; + } + Sleep(TDuration::MilliSeconds(200)); + } + + auto status = client.ResetOffset(path, TEST_CONSUMER, TResetOffsetSettings().Earliest()).GetValueSync(); + UNIT_ASSERT_C(status.IsSuccess(), status.GetIssues().ToString()); + + auto descr = setup.DescribeConsumer(TEST_TOPIC, TEST_CONSUMER); + bool sawInactive = false; + for (const auto& part : descr.GetPartitions()) { + if (!part.GetActive()) { + sawInactive = true; + const auto& stats = part.GetPartitionConsumerStats(); + UNIT_ASSERT(stats); + UNIT_ASSERT_VALUES_EQUAL(stats->GetCommittedOffset(), 0); + } + } + UNIT_ASSERT(sawInactive); +} + +} // TResetOffsetSdkTests diff --git a/src/client/types/credentials/CMakeLists.txt b/src/client/types/credentials/CMakeLists.txt index de608b39882..283dfe0e1fc 100644 --- a/src/client/types/credentials/CMakeLists.txt +++ b/src/client/types/credentials/CMakeLists.txt @@ -1,5 +1,6 @@ add_subdirectory(login) add_subdirectory(oauth2_token_exchange) +add_subdirectory(oidc) _ydb_sdk_add_library(client-types-credentials) @@ -10,6 +11,7 @@ target_link_libraries(client-types-credentials PUBLIC yql-public-issue client-types-credentials-login client-types-credentials-oauth2 + client-types-credentials-oidc ) target_sources(client-types-credentials PRIVATE diff --git a/src/client/types/credentials/oidc/CMakeLists.txt b/src/client/types/credentials/oidc/CMakeLists.txt new file mode 100644 index 00000000000..41c31415304 --- /dev/null +++ b/src/client/types/credentials/oidc/CMakeLists.txt @@ -0,0 +1,24 @@ +_ydb_sdk_add_library(client-types-credentials-oidc) + +target_sources(client-types-credentials-oidc PRIVATE + client_provider.cpp + credentials.cpp + device_provider.cpp + private.cpp + protocol.cpp + provider_base.cpp + static_provider.cpp +) + +target_link_libraries(client-types-credentials-oidc PUBLIC + yutil + http-simple + json + openssl-crypto + string_utils-base64 + string_utils-quote + client-types + client-types-credentials +) + +_ydb_sdk_install_targets(TARGETS client-types-credentials-oidc) diff --git a/src/client/types/credentials/oidc/client_provider.cpp b/src/client/types/credentials/oidc/client_provider.cpp new file mode 100644 index 00000000000..65c379e7c43 --- /dev/null +++ b/src/client/types/credentials/oidc/client_provider.cpp @@ -0,0 +1,19 @@ +#include "client_provider.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +TClientProvider::TClientProvider(const TOidcConfig& config, std::weak_ptr facility) + : TRefreshingProviderBase(config, std::move(facility)) +{ + Start(); +} + +TClientProvider::~TClientProvider() { + Stop(); +} + +TTokenCache TClientProvider::AcquireToken() { + return GetProtocol().ClientGrant(); +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/client_provider.h b/src/client/types/credentials/oidc/client_provider.h new file mode 100644 index 00000000000..7a593c68ff8 --- /dev/null +++ b/src/client/types/credentials/oidc/client_provider.h @@ -0,0 +1,16 @@ +#pragma once + +#include "provider_base.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +class TClientProvider final: public TRefreshingProviderBase { +public: + TClientProvider(const TOidcConfig& config, std::weak_ptr facility); + ~TClientProvider() override; + +private: + TTokenCache AcquireToken() override; +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/credentials.cpp b/src/client/types/credentials/oidc/credentials.cpp new file mode 100644 index 00000000000..8409f4fc2f5 --- /dev/null +++ b/src/client/types/credentials/oidc/credentials.cpp @@ -0,0 +1,102 @@ +#include +#include +#include +#include +#include + +#include +#include + +#include +#include + +namespace NYdb::inline V3::NOidc { + +ITokenCacher::~ITokenCacher() = default; + +IAuthAcceptor::~IAuthAcceptor() = default; + +bool TOAuthToken::IsValid(TInstant now) const { + return !Token.empty() && (!ExpiresAt.has_value() || ExpiresAt.value() > now); +} + +namespace NPrivate { + +namespace { + +class TFactory final: public ICredentialsProviderFactory { +public: + explicit TFactory(TOidcConfig config); + + TCredentialsProviderPtr CreateProvider() const override; + + TCredentialsProviderPtr CreateProvider(std::weak_ptr facility) const override; + + std::string GetClientIdentity() const override; + +private: + TCredentialsProviderPtr CreateProviderImpl(std::weak_ptr facility) const; + + TOidcConfig Config; + std::string Identity; + mutable TMutex Mutex; + mutable TCredentialsProviderPtr Provider; +}; + +TFactory::TFactory(TOidcConfig config) + : Config(std::move(config)) + , Identity(GetOidcClientIdentity(Config)) +{ + // The factory is identified before authorization, so a token's sub claim + // is not available for client/device grants. Keep the credential fingerprint + // stable and distinguish custom hooks by their process-local instance identity. + if (Config.Cacher_ != nullptr || Config.Acceptor_ != nullptr) { + Identity = HashIdentity(Identity + ":" + + std::to_string(reinterpret_cast(Config.Cacher_.get())) + ":" + + std::to_string(reinterpret_cast(Config.Acceptor_.get()))); + } +} + +TCredentialsProviderPtr TFactory::CreateProvider() const { + with_lock (Mutex) { + if (Provider == nullptr) { + auto facility = CreateSimpleCoreFacility(); + Provider = std::make_shared( + facility, CreateProviderImpl(facility)); + } + return Provider; + } +} + +TCredentialsProviderPtr TFactory::CreateProvider(std::weak_ptr facility) const { + return CreateProviderImpl(std::move(facility)); +} + +std::string TFactory::GetClientIdentity() const { + return Identity; +} + +TCredentialsProviderPtr TFactory::CreateProviderImpl(std::weak_ptr facility) const { + return std::visit(TOverloaded{ + [&](const TStaticOidcConfig&) -> TCredentialsProviderPtr { + return std::make_shared(Config, std::move(facility)); + }, + [&](const TClientOidcConfig&) -> TCredentialsProviderPtr { + return std::make_shared(Config, std::move(facility)); + }, + [&](const TDeviceOidcConfig&) -> TCredentialsProviderPtr { + return std::make_shared(Config, std::move(facility)); + }, + }, Config.FlowConfig); +} + +} // namespace + +} // namespace NPrivate + +std::shared_ptr CreateOidcProviderFactory(const TOidcConfig& config) { + NPrivate::ValidateOidcConfig(config); + return std::make_shared(config); +} + +} // namespace NYdb::inline V3::NOidc diff --git a/src/client/types/credentials/oidc/device_provider.cpp b/src/client/types/credentials/oidc/device_provider.cpp new file mode 100644 index 00000000000..fbb9b1f00f2 --- /dev/null +++ b/src/client/types/credentials/oidc/device_provider.cpp @@ -0,0 +1,19 @@ +#include "device_provider.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +TDeviceProvider::TDeviceProvider(const TOidcConfig& config, std::weak_ptr facility) + : TRefreshingProviderBase(config, std::move(facility)) +{ + Start(); +} + +TDeviceProvider::~TDeviceProvider() { + Stop(); +} + +TTokenCache TDeviceProvider::AcquireToken() { + return GetProtocol().DeviceGrant([this](TDuration delay) { return Wait(delay); }); +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/device_provider.h b/src/client/types/credentials/oidc/device_provider.h new file mode 100644 index 00000000000..6d9fa02a0ce --- /dev/null +++ b/src/client/types/credentials/oidc/device_provider.h @@ -0,0 +1,16 @@ +#pragma once + +#include "provider_base.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +class TDeviceProvider final: public TRefreshingProviderBase { +public: + TDeviceProvider(const TOidcConfig& config, std::weak_ptr facility); + ~TDeviceProvider() override; + +private: + TTokenCache AcquireToken() override; +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/private.cpp b/src/client/types/credentials/oidc/private.cpp new file mode 100644 index 00000000000..989ca86f083 --- /dev/null +++ b/src/client/types/credentials/oidc/private.cpp @@ -0,0 +1,227 @@ +#include "private.h" + +#include +#include +#include + +#include + +#include +#include +#include + +namespace NYdb::inline V3::NOidc::NPrivate { +namespace { + +constexpr size_t MaxJwtPayloadSize = 1024 * 1024; +// Seconds are converted to signed chrono microseconds. Reserve half the range +// for deadline arithmetic and doubling the polling interval without overflow. +constexpr ui64 MaxDurationSeconds = std::numeric_limits::max() / 2 / 1'000'000; + +} // namespace + +std::string HashIdentity(const std::string& data) { + const auto hash = NOpenSsl::NSha256::Calc(data.data(), data.size()); + static constexpr char Hex[] = "0123456789abcdef"; + std::string identity = "oidc:"; + for (const auto byte : hash) { + identity += Hex[byte >> 4]; + identity += Hex[byte & 0xF]; + } + return identity; +} + +void ValidateOidcConfig(const TOidcConfig& config) { + ParseUrl(config.Issuer, true); + std::visit(TOverloaded{ + [](const TStaticOidcConfig& flow) { + if (flow.AccessToken.empty()) { + throw std::invalid_argument("OIDC credentials: static_credentials requires access_token"); + } + }, + [](const TClientOidcConfig& flow) { + if (flow.ClientId.empty()) { + throw std::invalid_argument("OIDC credentials: client_id is required"); + } + if (flow.ClientSecret.empty()) { + throw std::invalid_argument("OIDC credentials: client_secret is required"); + } + }, + [](const TDeviceOidcConfig& flow) { + if (flow.ClientId.empty()) { + throw std::invalid_argument("OIDC credentials: client_id is required"); + } + }, + }, config.FlowConfig); + for (const auto& scope : Scopes(config)) { + if (scope.empty() || std::any_of(scope.begin(), scope.end(), [](unsigned char c) { + return c <= 0x20 || c >= 0x7f || c == '"' || c == '\\'; + })) { + throw std::invalid_argument("OIDC credentials: invalid scope"); + } + } +} + +std::string GetOidcClientIdentity(const TOidcConfig& config) { + std::string data; + const auto append = [&data](const std::string& value) { + data += std::to_string(value.size()) + ":" + value; + }; + append(config.Issuer); + append(std::to_string(config.FlowConfig.index())); + append(ClientId(config)); + append(ClientSecret(config)); + auto scopes = Scopes(config); + std::sort(scopes.begin(), scopes.end()); + scopes.erase(std::unique(scopes.begin(), scopes.end()), scopes.end()); + for (const auto& scope : scopes) { + append(scope); + } + if (const auto* flow = std::get_if(&config.FlowConfig); flow != nullptr) { + append(flow->AccessToken); + append(flow->ExpiresAt.has_value() ? std::to_string(flow->ExpiresAt->MicroSeconds()) : ""); + } + return HashIdentity(data); +} + +TError::TError(const std::string& message, bool retryable, std::string code) + : std::runtime_error("OIDC credentials: " + message) + , Retryable(retryable) + , Code(std::move(code)) +{ +} + +const NJson::TJsonValue* Field(const NJson::TJsonValue& json, const TString& name) { + const auto& map = json.GetMapSafe(); + const auto it = map.find(name); + return (it == map.end()) ? nullptr : &it->second; +} + +ui64 Seconds(const NJson::TJsonValue& value, const std::string& field, bool allowZero) { + if ((!value.IsInteger() && !value.IsUInteger()) || (value.IsInteger() && value.GetInteger() < 0)) { + throw TError("invalid " + field, false, {}); + } + const auto seconds = value.GetUInteger(); + if ((!seconds && !allowZero) || seconds > MaxDurationSeconds) { + throw TError("invalid " + field, false, {}); + } + return seconds; +} + +NUri::TUri ParseUrl(const std::string& value, bool issuer) { + const std::string role = issuer ? "issuer" : "endpoint"; + const auto invalidUrl = [&role] { + return std::invalid_argument("OIDC credentials: invalid " + role + " URL"); + }; + if (value.empty()) { + throw invalidUrl(); + } + const bool hasControlCharacters = std::any_of(value.begin(), value.end(), [](unsigned char c) { + return c <= 0x20 || c == 0x7f; + }); + if (hasControlCharacters) { + throw invalidUrl(); + } + + NUri::TUri url; + if (url.Parse(value, NUri::TFeature::FeaturesAll) != NUri::TUri::TState::EParsed::ParsedOK) { + throw invalidUrl(); + } + if (url.GetHost().empty()) { + throw invalidUrl(); + } + for (const auto field : {NUri::TUri::FieldUser, NUri::TUri::FieldPass, NUri::TUri::FieldFrag}) { + if (!url.IsNull(field)) { + throw invalidUrl(); + } + } + if (issuer && !url.IsNull(NUri::TUri::FieldQuery)) { + throw invalidUrl(); + } + if (url.GetScheme() != NUri::TScheme::SchemeHTTPS) { + throw std::invalid_argument("OIDC credentials: " + role + " requires HTTPS"); + } + if (!url.GetPort()) { + throw std::invalid_argument("OIDC credentials: invalid " + role + " port"); + } + return url; +} + +std::optional JwtExpiry(const std::string& token) { + const auto first = token.find('.'); + const auto second = (first == std::string::npos) ? first : token.find('.', first + 1); + if (second == std::string::npos || second - first > MaxJwtPayloadSize) { + return std::nullopt; + } + NJson::TJsonValue payload; + try { + const auto decoded = Base64DecodeUneven(TStringBuf(token.data() + first + 1, second - first - 1)); + if (!NJson::ReadJsonTree(decoded, &payload) || !payload.IsMap()) { + return std::nullopt; + } + } catch (const std::exception&) { + return std::nullopt; + } + if (const auto* expiry = Field(payload, "exp"); expiry != nullptr) { + // A negative NumericDate is in the past, not an unknown lifetime. + if (expiry->IsInteger() && expiry->GetInteger() < 0) { + return TInstant::Zero(); + } + try { + return TInstant::Seconds(Seconds(*expiry, "exp", true)); + } catch (const TError&) { + // JWT decoding is only a scheduling hint; token validation belongs + // to the server, just as it does for opaque access tokens. + return std::nullopt; + } + } + return std::nullopt; +} + +TDuration TDevicePolling::NextDelay(TInstant now) const { + if (now >= Deadline) { + throw TError("device authorization expired", false, "expired_token"); + } + return std::min(Interval, Deadline - now); +} + +void TDevicePolling::HandleError(const TError& error) { + if (error.Code == "authorization_pending") { + return; + } + if (error.Code == "slow_down") { + Interval = std::min(Interval + TDuration::Seconds(5), TDuration::Hours(24)); + return; + } + if (error.Retryable) { + Interval = std::min(Interval * 2, TDuration::Hours(24)); + return; + } + throw error; +} + +std::string ClientId(const TOidcConfig& config) { + return std::visit(TOverloaded{ + [](const TStaticOidcConfig&) { return std::string{}; }, + [](const TClientOidcConfig& flow) { return flow.ClientId; }, + [](const TDeviceOidcConfig& flow) { return flow.ClientId; }, + }, config.FlowConfig); +} + +std::string ClientSecret(const TOidcConfig& config) { + return std::visit(TOverloaded{ + [](const TStaticOidcConfig&) { return std::string{}; }, + [](const TClientOidcConfig& flow) { return flow.ClientSecret; }, + [](const TDeviceOidcConfig&) { return std::string{}; }, + }, config.FlowConfig); +} + +std::vector Scopes(const TOidcConfig& config) { + return std::visit(TOverloaded{ + [](const TStaticOidcConfig&) { return std::vector{}; }, + [](const TClientOidcConfig& flow) { return flow.Scopes; }, + [](const TDeviceOidcConfig& flow) { return flow.Scopes; }, + }, config.FlowConfig); +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/private.h b/src/client/types/credentials/oidc/private.h new file mode 100644 index 00000000000..4278023bcf2 --- /dev/null +++ b/src/client/types/credentials/oidc/private.h @@ -0,0 +1,43 @@ +#pragma once + +#include + +#include +#include + +#include + +namespace NYdb::inline V3::NOidc::NPrivate { + +void ValidateOidcConfig(const TOidcConfig& config); +std::string HashIdentity(const std::string& data); + +// Deterministic credential fingerprint; excludes cacher/acceptor instances. +std::string GetOidcClientIdentity(const TOidcConfig& config); + +class TError: public std::runtime_error { +public: + explicit TError(const std::string& message, bool retryable, std::string code); + + bool Retryable; + std::string Code; +}; + +const NJson::TJsonValue* Field(const NJson::TJsonValue& json, const TString& name); +ui64 Seconds(const NJson::TJsonValue& value, const std::string& field, bool allowZero); + +NUri::TUri ParseUrl(const std::string& value, bool issuer); +std::optional JwtExpiry(const std::string& token); +std::string ClientId(const TOidcConfig& config); +std::string ClientSecret(const TOidcConfig& config); +std::vector Scopes(const TOidcConfig& config); + +struct TDevicePolling { + TDuration Interval = TDuration::Seconds(5); + TInstant Deadline; + + TDuration NextDelay(TInstant now) const; + void HandleError(const TError& error); +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/protocol.cpp b/src/client/types/credentials/oidc/protocol.cpp new file mode 100644 index 00000000000..26170bb613d --- /dev/null +++ b/src/client/types/credentials/oidc/protocol.cpp @@ -0,0 +1,326 @@ +#include "protocol.h" + +#include + +#include +#include +#include +#include + +#include +#include +#include + +#include + +namespace NYdb::inline V3::NOidc::NPrivate { +namespace { + +constexpr size_t MaxResponseSize = 1024 * 1024; +const TDuration SocketTimeout = TDuration::Seconds(5); +const TDuration ConnectTimeout = TDuration::Seconds(30); + +class TResponseBuffer: public IOutputStream { +public: + TString Body; + +private: + void DoWrite(const void* buffer, size_t size) override; +}; + +std::string String(const NJson::TJsonValue& json, const TString& name, bool required); +std::string ScopeString(const TOidcConfig& config); +bool IsBearer(const std::string& value); +void CheckToken(const std::string& token); + +void TResponseBuffer::DoWrite(const void* buffer, size_t size) { + if (size > MaxResponseSize - Body.size()) { + throw TError("response exceeds size limit", false, {}); + } + Body.append(static_cast(buffer), size); +} + +std::string String(const NJson::TJsonValue& json, const TString& name, bool required) { + const auto* field = Field(json, name); + if (field == nullptr && !required) { + return {}; + } + if (field == nullptr || !field->IsString() || field->GetString().empty()) { + throw TError("missing or invalid " + std::string(name), false, {}); + } + return std::string(field->GetString()); +} + +std::string ScopeString(const TOidcConfig& config) { + auto scopes = Scopes(config); + if (std::find(scopes.begin(), scopes.end(), "openid") == scopes.end()) { + scopes.push_back("openid"); + } + std::string result; + for (const auto& scope : scopes) { + if (!result.empty()) { + result += ' '; + } + result += scope; + } + return result; +} + +bool IsBearer(const std::string& value) { + return to_lower(TString(value)) == "bearer"; +} + +void CheckToken(const std::string& token) { + if (std::any_of(token.begin(), token.end(), [](unsigned char c) { return c <= 0x20 || c >= 0x7f; })) { + throw TError("invalid token characters", false, {}); + } +} + +} // namespace + +TProtocol::TProtocol(const TOidcConfig& config, NThreading::TCancellationToken cancellation) + : Config(config) + , Cancellation(std::move(cancellation)) +{ +} + +TProtocol::~TProtocol() = default; + +NJson::TJsonValue TProtocol::Request(const std::string& endpoint, const TCgiParameters* form, bool authenticate, TInstant deadline) { + Cancellation.ThrowIfCancellationRequested(); + const auto url = ParseUrl(endpoint, false); + TKeepAliveHttpClient::THeaders headers; + TCgiParameters body = (form != nullptr) ? *form : TCgiParameters{}; + if (authenticate) { + TString secret(ClientSecret(Config)); + if (!secret.empty()) { + TString clientId(ClientId(Config)); + Quote(clientId, ""); + Quote(secret, ""); + headers["Authorization"] = "Basic " + Base64Encode(clientId + ":" + secret); + } else { + body.InsertUnescaped("client_id", ClientId(Config)); + } + } + headers["Accept"] = "application/json"; + if (form != nullptr) { + headers["Content-Type"] = "application/x-www-form-urlencoded"; + } + const auto now = TInstant::Now(); + if (deadline <= now) { + throw TError("HTTP request deadline exceeded", true, {}); + } + const auto remaining = deadline - now; + const auto host = url.PrintS(NUri::TUri::FlagScheme | NUri::TUri::FlagHost | NUri::TUri::FlagHostAscii); + const auto path = url.PrintS(NUri::TUri::FlagPath | NUri::TUri::FlagQuery); + const auto encodedBody = body.Print(); + TStringBuilder request; + request << (form != nullptr ? "POST " : "GET ") << path << " HTTP/1.1\r\n" + << "Host: " << url.PrintS(NUri::TUri::FlagHost | NUri::TUri::FlagHostAscii | NUri::TUri::FlagPort) << "\r\n" + << "Content-Length: " << encodedBody.size() << "\r\n"; + for (const auto& [name, value] : headers) { + request << name << ": " << value << "\r\n"; + } + request << "\r\n" << encodedBody; + TResponseBuffer response; + unsigned status; + try { + // Shutdown waits for this synchronous request. Keep HTTP cancellation + // subscriptions local to the request instead of retaining them until shutdown. + NThreading::TCancellationTokenSource requestCancellation; + TKeepAliveHttpClient client(host, url.GetPort(), + std::min(SocketTimeout, remaining), std::min(ConnectTimeout, remaining), false, false, true); + status = client.DoRequestRaw(request, &response, nullptr, requestCancellation.Token()); + Cancellation.ThrowIfCancellationRequested(); + } catch (const TError&) { + throw; + } catch (const std::exception&) { + Cancellation.ThrowIfCancellationRequested(); + throw TError("HTTP transport failed", true, {}); + } + if (TInstant::Now() >= deadline) { + throw TError("HTTP request deadline exceeded", true, {}); + } + NJson::TJsonValue json; + NJson::TJsonReaderConfig reader; + reader.MaxDepth = 32; + const bool valid = NJson::ReadJsonTree(response.Body, &reader, &json) && json.IsMap(); + if (status != 200) { + std::string code; + if (valid) { + if (const auto* error = Field(json, "error"); error != nullptr && error->IsString()) { + for (const char* known : {"invalid_grant", "invalid_client", "invalid_scope", "unauthorized_client", + "unsupported_grant_type", "authorization_pending", "slow_down", "access_denied", "expired_token"}) { + if (error->GetString() == known) { + code = known; + break; + } + } + } + } + throw TError("HTTP " + std::to_string(status) + (code.empty() ? "" : " (" + code + ")"), + status == 408 || status == 429 || status == 500 || status == 502 || status == 503 || status == 504, code); + } + if (!valid) { + throw TError("invalid JSON response", false, {}); + } + return json; +} + +void TProtocol::Discover() { + if (!TokenEndpoint.empty()) { + return; + } + auto discovery = Config.Issuer; + while (!discovery.empty() && discovery.back() == '/') { + discovery.pop_back(); + } + const auto metadata = Request(discovery + "/.well-known/openid-configuration", nullptr, false, TInstant::Max()); + // Only the discovery request path is normalized; the issuer identifier + // must match exactly (OpenID Connect Discovery 1.0, sections 4.1 and 4.3). + const auto advertisedIssuer = String(metadata, "issuer", true); + if (advertisedIssuer != Config.Issuer) { + try { + // Reject userinfo, queries and control characters before including + // an untrusted metadata value in diagnostics. + ParseUrl(advertisedIssuer, true); + } catch (const std::invalid_argument&) { + throw TError("discovery issuer mismatch: invalid advertised issuer URL", false, {}); + } + throw TError("discovery issuer mismatch: configured '" + Config.Issuer + + "', advertised '" + advertisedIssuer + "'", false, {}); + } + auto tokenEndpoint = String(metadata, "token_endpoint", true); + auto deviceEndpoint = String(metadata, "device_authorization_endpoint", false); + ParseUrl(tokenEndpoint, false); + if (!deviceEndpoint.empty()) { + ParseUrl(deviceEndpoint, false); + } + if (!ClientSecret(Config).empty()) { + if (const auto* methods = Field(metadata, "token_endpoint_auth_methods_supported"); methods != nullptr) { + if (!methods->IsArray()) { + throw TError("invalid token_endpoint_auth_methods_supported", false, {}); + } + bool supported = false; + for (const auto& method : methods->GetArray()) { + if (!method.IsString()) { + throw TError("invalid token_endpoint_auth_methods_supported", false, {}); + } + supported |= (method.GetString() == "client_secret_basic"); + } + if (!supported) { + throw TError("client_secret_basic is not supported", false, {}); + } + } + } + TokenEndpoint = std::move(tokenEndpoint); + DeviceEndpoint = std::move(deviceEndpoint); +} + +TTokenCache TProtocol::TokenRequest(TCgiParameters form, const std::optional& refresh, TInstant deadline) { + Discover(); + const auto now = TInstant::Now(); + const auto response = Request(TokenEndpoint, &form, true, deadline); + // Include JSON response processing in the device authorization deadline. + if (TInstant::Now() >= deadline) { + throw TError("device authorization expired", false, "expired_token"); + } + if (!IsBearer(String(response, "token_type", true))) { + throw TError("unsupported token_type", false, {}); + } + TTokenCache result; + result.AccessToken.Token = String(response, "access_token", true); + CheckToken(result.AccessToken.Token); + if (const auto* expires = Field(response, "expires_in"); expires != nullptr) { + result.AccessToken.ExpiresAt = now + TDuration::Seconds(Seconds(*expires, "expires_in", false)); + } else { + result.AccessToken.ExpiresAt = JwtExpiry(result.AccessToken.Token); + } + if (!result.AccessToken.IsValid(TInstant::Now())) { + throw TError("access token expired", false, {}); + } + const auto refreshToken = String(response, "refresh_token", false); + if (!refreshToken.empty()) { + CheckToken(refreshToken); + result.RefreshToken = TOAuthToken{refreshToken, std::nullopt}; + } else { + result.RefreshToken = refresh; + } + if (result.RefreshToken.has_value()) { + if (const auto* expires = Field(response, "refresh_expires_in"); expires != nullptr) { + const auto seconds = Seconds(*expires, "refresh_expires_in", true); + result.RefreshToken->ExpiresAt = seconds + ? std::optional(now + TDuration::Seconds(seconds)) + : std::nullopt; + } + } + return result; +} + +TTokenCache TProtocol::Refresh(const TOAuthToken& refresh) { + TCgiParameters form; + form.InsertUnescaped("grant_type", "refresh_token"); + form.InsertUnescaped("refresh_token", refresh.Token); + return TokenRequest(std::move(form), refresh, TInstant::Max()); +} + +TTokenCache TProtocol::ClientGrant() { + TCgiParameters form; + form.InsertUnescaped("grant_type", "client_credentials"); + form.InsertUnescaped("scope", ScopeString(Config)); + return TokenRequest(std::move(form), std::nullopt, TInstant::Max()); +} + +TTokenCache TProtocol::DeviceGrant(const std::function& wait) { + if (Config.Acceptor_ == nullptr) { + throw TError("device authorization requires an auth acceptor", false, {}); + } + Discover(); + if (DeviceEndpoint.empty()) { + throw TError("discovery is missing device_authorization_endpoint", false, {}); + } + TCgiParameters form; + form.InsertUnescaped("client_id", ClientId(Config)); + form.InsertUnescaped("scope", ScopeString(Config)); + const auto started = TInstant::Now(); + const auto response = Request(DeviceEndpoint, &form, false, TInstant::Max()); + TDeviceAuthInfo info; + info.UserCode = String(response, "user_code", true); + info.VerificationUrl = String(response, "verification_uri", true); + ParseUrl(info.VerificationUrl, false); + const auto complete = String(response, "verification_uri_complete", false); + if (!complete.empty()) { + ParseUrl(complete, false); + info.VerificationUrlComplete = complete; + } + const auto deviceCode = String(response, "device_code", true); + const auto* expires = Field(response, "expires_in"); + if (expires == nullptr) { + throw TError("missing device expires_in", false, {}); + } + info.ExpiresAt = started + TDuration::Seconds(Seconds(*expires, "expires_in", false)); + TDevicePolling polling; + polling.Deadline = info.ExpiresAt; + if (const auto* interval = Field(response, "interval"); interval != nullptr) { + polling.Interval = TDuration::Seconds(Seconds(*interval, "interval", false)); + } + Config.Acceptor_->Accept(info); + TCgiParameters tokenForm; + tokenForm.InsertUnescaped("grant_type", "urn:ietf:params:oauth:grant-type:device_code"); + tokenForm.InsertUnescaped("device_code", deviceCode); + for (;;) { + if (!wait(polling.NextDelay(TInstant::Now()))) { + throw TError("provider stopped", false, {}); + } + // Do not start another token request if the wait overshot the device + // deadline. This also preserves expired_token instead of an HTTP timeout. + polling.NextDelay(TInstant::Now()); + try { + return TokenRequest(tokenForm, std::nullopt, polling.Deadline); + } catch (const TError& error) { + polling.HandleError(error); + } + } +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/protocol.h b/src/client/types/credentials/oidc/protocol.h new file mode 100644 index 00000000000..0ab4d3f9630 --- /dev/null +++ b/src/client/types/credentials/oidc/protocol.h @@ -0,0 +1,33 @@ +#pragma once + +#include + +#include +#include +#include + +#include + +namespace NYdb::inline V3::NOidc::NPrivate { + +class TProtocol { +public: + TProtocol(const TOidcConfig& config, NThreading::TCancellationToken cancellation); + ~TProtocol(); + + TTokenCache Refresh(const TOAuthToken& refresh); + TTokenCache ClientGrant(); + TTokenCache DeviceGrant(const std::function& wait); + +private: + NJson::TJsonValue Request(const std::string& endpoint, const TCgiParameters* form, bool authenticate, TInstant deadline); + void Discover(); + TTokenCache TokenRequest(TCgiParameters form, const std::optional& refresh, TInstant deadline); + + const TOidcConfig& Config; + NThreading::TCancellationToken Cancellation; + std::string TokenEndpoint; + std::string DeviceEndpoint; +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/provider_base.cpp b/src/client/types/credentials/oidc/provider_base.cpp new file mode 100644 index 00000000000..7b32c57d45e --- /dev/null +++ b/src/client/types/credentials/oidc/provider_base.cpp @@ -0,0 +1,350 @@ +#include "provider_base.h" + +#include + +#include + +namespace NYdb::inline V3::NOidc::NPrivate { +namespace { + +std::exception_ptr StoppedError(); +void SetException(NThreading::TPromise promise, std::exception_ptr error) noexcept; + +std::exception_ptr StoppedError() { + return std::make_exception_ptr(TError("provider stopped", false, {})); +} + +void SetException(NThreading::TPromise promise, std::exception_ptr error) noexcept { + try { + promise.TrySetException(std::move(error)); + } catch (...) { + // The promise is already settled; a throwing subscriber must not interrupt cleanup. + } +} + +} // namespace + +TProviderBase::TProviderBase(TOidcConfig config, std::weak_ptr facility) + : Config(std::move(config)) + , Facility(std::move(facility)) + , Pending(NThreading::NewPromise()) +{ +} + +TProviderBase::~TProviderBase() { + Stop(); +} + +void TProviderBase::Start() { + try { + Worker = std::thread([this] { Run(); }); + } catch (...) { + Fail(std::current_exception()); + } +} + +std::string TProviderBase::GetAuthInfo() const { + return GetAuthInfoAsync().GetValueSync(); +} + +NThreading::TFuture TProviderBase::GetAuthInfoAsync() const { + std::string token; + std::exception_ptr error; + with_lock (Mutex) { + if (Stopping || Facility.expired()) { + error = StoppedError(); + } else if (Tokens.has_value() && Tokens->AccessToken.IsValid(TInstant::Now())) { + token = "Bearer " + Tokens->AccessToken.Token; + } else if (Error != nullptr) { + error = Error; + } else { + return Pending.GetFuture(); + } + } + if (error != nullptr) { + return NThreading::MakeErrorFuture(error); + } + return NThreading::MakeFuture(std::move(token)); +} + +bool TProviderBase::IsValid() const { + with_lock (Mutex) { + return !Stopping && !Facility.expired() && + ((Tokens.has_value() && Tokens->AccessToken.IsValid(TInstant::Now())) || Error == nullptr); + } +} + +void TProviderBase::Stop() { + RequestStop(); + if (Worker.joinable()) { + Worker.join(); + } + CancelDeliveries(); +} + +void TProviderBase::RequestStop() { + NThreading::TPromise pending; + with_lock (Mutex) { + if (Stopping) { + return; + } + Stopping = true; + pending = Pending; + } + Changed.notify_all(); + Cancellation.Cancel(); + SetException(pending, StoppedError()); + CancelDeliveries(); +} + +void TProviderBase::Run() { + try { + if (IsStopped()) { + RequestStop(); + return; + } + RunTokens(); + } catch (...) { + Fail(std::current_exception()); + } + + for (;;) { + CompleteDiscardedDeliveries(); + with_lock (Mutex) { + const bool finished = std::all_of(Deliveries.begin(), Deliveries.end(), [](const auto& delivery) { + return delivery.Promise.GetFuture().HasValue() || delivery.Promise.GetFuture().HasException(); + }); + if (finished) { + return; + } + } + if (!Wait(TDuration::MilliSeconds(100))) { + return; + } + } +} + +TRefreshingProviderBase::TRefreshingProviderBase(const TOidcConfig& config, std::weak_ptr facility) + : TProviderBase(config, std::move(facility)) + , Protocol(Config, Cancellation.Token()) +{ +} + +void TRefreshingProviderBase::RunTokens() { + TTokenCache current = ReadCache().value_or(TTokenCache{}); + + const bool unknownRefresh = !current.AccessToken.ExpiresAt.has_value() && current.RefreshToken.has_value(); + if (current.AccessToken.IsValid(TInstant::Now()) && !unknownRefresh) { + Publish(current); + if (!WaitForRefresh(current)) { + return; + } + } + + TDuration retryDelay = TDuration::MilliSeconds(200); + for (;;) { + if (IsStopped()) { + RequestStop(); + return; + } + try { + current = Update(current); + Write(current); + Publish(current); + retryDelay = TDuration::MilliSeconds(200); + if (!WaitForRefresh(current)) { + return; + } + } catch (const TError& error) { + if (!error.Retryable) { + throw; + } + // Settle existing waiters while retrying in the background. GetAuthInfoAsync() + // keeps serving a still-valid token; Publish() clears the error on recovery. + Fail(std::current_exception()); + if (!Wait(retryDelay)) { + return; + } + retryDelay = std::min(retryDelay * 2, TDuration::Seconds(30)); + } + } +} + +bool TProviderBase::IsStopped() const { + with_lock (Mutex) { + return Stopping || Facility.expired(); + } +} + +bool TProviderBase::Wait(TDuration delay) { + with_lock (Mutex) { + auto remaining = std::chrono::microseconds(delay.MicroSeconds()); + const auto end = std::chrono::steady_clock::now() + remaining; + while (!Stopping && !Facility.expired()) { + { + auto unguard = Unguard(Mutex); + CompleteDiscardedDeliveries(); + } + if (Stopping) { + break; + } + if (remaining <= std::chrono::microseconds::zero()) { + return true; + } + Changed.wait_for(Mutex, std::min(remaining, std::chrono::microseconds(100'000)), [this] { return Stopping; }); + remaining = std::chrono::duration_cast(end - std::chrono::steady_clock::now()); + } + } + RequestStop(); + return false; +} + +bool TRefreshingProviderBase::WaitForRefresh(const TTokenCache& current) { + if (!current.AccessToken.ExpiresAt.has_value()) { + return false; + } + const auto now = TInstant::Now(); + if (*current.AccessToken.ExpiresAt <= now) { + return true; + } + // Preserve short lifetimes: a fixed multi-second floor could delay refresh + // until after expiry. Never reinterpret a known expiry as an unknown one. + return Wait(std::max((*current.AccessToken.ExpiresAt - now) / 2, TDuration::MilliSeconds(1))); +} + +void TProviderBase::Write(const TTokenCache& tokens) const { + try { + if (Config.Cacher_ != nullptr) { + Config.Cacher_->Write(tokens); + } + } catch (...) { + // Persistence is optional; keep the acquired token usable in memory. + // User-supplied exception messages may contain credentials. + } +} + +TTokenCache TRefreshingProviderBase::Update(const TTokenCache& current) { + if (current.RefreshToken.has_value() && current.RefreshToken->IsValid(TInstant::Now())) { + try { + return Protocol.Refresh(*current.RefreshToken); + } catch (const TError& error) { + if (error.Retryable || error.Code != "invalid_grant") { + throw; + } + } + } + return AcquireToken(); +} + +void TProviderBase::Complete(NThreading::TPromise pending, std::optional token, std::exception_ptr error) { + const auto callbackLifetime = std::make_shared(0); + bool stopped; + with_lock (Mutex) { + stopped = Stopping; + Deliveries.erase(std::remove_if(Deliveries.begin(), Deliveries.end(), [](const auto& delivery) { + return delivery.Promise.GetFuture().HasValue() || delivery.Promise.GetFuture().HasException(); + }), Deliveries.end()); + if (!stopped) { + Deliveries.push_back({pending, callbackLifetime}); + } + } + if (stopped) { + SetException(pending, StoppedError()); + return; + } + auto completion = [pending, token = std::move(token), error, callbackLifetime]() mutable { + Y_UNUSED(callbackLifetime); + try { + if (error != nullptr) { + SetException(pending, error); + } else if (!token->IsValid(TInstant::Now())) { + SetException(pending, std::make_exception_ptr(TError("access token expired before delivery", false, {}))); + } else { + pending.TrySetValue("Bearer " + token->Token); + } + } catch (...) { + // Covers both preparation failures and subscribers throwing after + // settlement. Neither may escape into the response queue executor. + SetException(pending, std::current_exception()); + } + }; + try { + if (auto facility = Facility.lock(); facility != nullptr) { + facility->PostToResponseQueue(std::move(completion)); + } else { + SetException(pending, StoppedError()); + } + } catch (...) { + SetException(pending, std::current_exception()); + } +} + +void TProviderBase::CompleteDiscardedDeliveries() { + std::vector> discarded; + with_lock (Mutex) { + auto it = Deliveries.begin(); + while (it != Deliveries.end()) { + const auto future = it->Promise.GetFuture(); + if (future.HasValue() || future.HasException()) { + it = Deliveries.erase(it); + } else if (it->CallbackLifetime.expired()) { + discarded.push_back(it->Promise); + it = Deliveries.erase(it); + } else { + ++it; + } + } + } + + for (auto& promise : discarded) { + SetException(promise, StoppedError()); + } +} + +void TProviderBase::Publish(const TTokenCache& current) { + NThreading::TPromise pending; + with_lock (Mutex) { + if (Stopping) { + return; + } + Tokens = current; + Error = nullptr; + pending = Pending; + Pending = NThreading::NewPromise(); + } + Complete(pending, current.AccessToken, {}); +} + +void TProviderBase::Fail(std::exception_ptr error) { + NThreading::TPromise pending; + with_lock (Mutex) { + Error = error; + pending = Pending; + } + Complete(pending, std::nullopt, error); +} + +void TProviderBase::CancelDeliveries() { + std::vector deliveries; + with_lock (Mutex) { + deliveries.swap(Deliveries); + } + for (auto& delivery : deliveries) { + SetException(delivery.Promise, StoppedError()); + } +} + +std::optional TProviderBase::ReadCache() const { + try { + return Config.Cacher_ != nullptr ? Config.Cacher_->Read() : std::nullopt; + } catch (...) { + // A broken cache must not prevent a fresh authorization attempt. + return std::nullopt; + } +} + +TProtocol& TRefreshingProviderBase::GetProtocol() { + return Protocol; +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/provider_base.h b/src/client/types/credentials/oidc/provider_base.h new file mode 100644 index 00000000000..47ea5b2321c --- /dev/null +++ b/src/client/types/credentials/oidc/provider_base.h @@ -0,0 +1,77 @@ +#pragma once + +#include "private.h" +#include "protocol.h" + +#include + +#include +#include + +namespace NYdb::inline V3::NOidc::NPrivate { + +class TProviderBase: public ICredentialsProvider { + struct TDelivery { + NThreading::TPromise Promise; + std::weak_ptr CallbackLifetime; + }; + +public: + TProviderBase(TOidcConfig config, std::weak_ptr facility); + ~TProviderBase() override; + + std::string GetAuthInfo() const override; + NThreading::TFuture GetAuthInfoAsync() const override; + bool IsValid() const override; + void Stop(); + +protected: + void Start(); + virtual void RunTokens() = 0; + + bool Wait(TDuration delay); + std::optional ReadCache() const; + void Write(const TTokenCache& tokens) const; + void Fail(std::exception_ptr error); + void Publish(const TTokenCache& current); + bool IsStopped() const; + void RequestStop(); + + TOidcConfig Config; + NThreading::TCancellationTokenSource Cancellation; + +private: + void Run(); + void CancelDeliveries(); + void Complete(NThreading::TPromise pending, std::optional token, std::exception_ptr error); + void CompleteDiscardedDeliveries(); + + std::weak_ptr Facility; + mutable TMutex Mutex; + std::condition_variable_any Changed; + bool Stopping = false; + std::optional Tokens; + std::exception_ptr Error; + NThreading::TPromise Pending; + std::vector Deliveries; + std::thread Worker; +}; + +class TRefreshingProviderBase: public TProviderBase { +public: + TRefreshingProviderBase(const TOidcConfig& config, std::weak_ptr facility); + +protected: + virtual TTokenCache AcquireToken() = 0; + + TProtocol& GetProtocol(); + +private: + void RunTokens() override; + bool WaitForRefresh(const TTokenCache& current); + TTokenCache Update(const TTokenCache& current); + + TProtocol Protocol; +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/static_provider.cpp b/src/client/types/credentials/oidc/static_provider.cpp new file mode 100644 index 00000000000..3fe6f69e02f --- /dev/null +++ b/src/client/types/credentials/oidc/static_provider.cpp @@ -0,0 +1,32 @@ +#include "static_provider.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +TStaticProvider::TStaticProvider(const TOidcConfig& config, std::weak_ptr facility) + : TProviderBase(config, std::move(facility)) +{ + Start(); +} + +TStaticProvider::~TStaticProvider() { + Stop(); +} + +void TStaticProvider::RunTokens() { + const auto& flow = std::get(Config.FlowConfig); + TTokenCache current; + current.AccessToken = {flow.AccessToken, flow.ExpiresAt}; + if (!current.AccessToken.ExpiresAt.has_value()) { + current.AccessToken.ExpiresAt = JwtExpiry(current.AccessToken.Token); + } + if (!current.AccessToken.IsValid(TInstant::Now())) { + throw TError("static credentials have expired", false, {}); + } + Write(current); + Publish(current); + if (current.AccessToken.ExpiresAt.has_value()) { + Fail(std::make_exception_ptr(TError("static credentials have expired", false, {}))); + } +} + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/client/types/credentials/oidc/static_provider.h b/src/client/types/credentials/oidc/static_provider.h new file mode 100644 index 00000000000..9e6450c0dab --- /dev/null +++ b/src/client/types/credentials/oidc/static_provider.h @@ -0,0 +1,16 @@ +#pragma once + +#include "provider_base.h" + +namespace NYdb::inline V3::NOidc::NPrivate { + +class TStaticProvider final: public TProviderBase { +public: + TStaticProvider(const TOidcConfig& config, std::weak_ptr facility); + ~TStaticProvider() override; + +private: + void RunTokens() override; +}; + +} // namespace NYdb::inline V3::NOidc::NPrivate diff --git a/src/library/grpc/client/grpc_client_low.h b/src/library/grpc/client/grpc_client_low.h index 42852001e8a..3d83dcc5cac 100644 --- a/src/library/grpc/client/grpc_client_low.h +++ b/src/library/grpc/client/grpc_client_low.h @@ -485,6 +485,22 @@ class IStreamRequestReadWriteProcessor : public IStreamRequestReadProcessor; using TReadCallback = typename TBase::TReadCallback; using TWriteCallback = typename TBase::TWriteCallback; - using TAsyncReaderWriterPtr = std::unique_ptr>; - using TAsyncRequest = TAsyncReaderWriterPtr (TStub::*)(grpc::ClientContext*, grpc::CompletionQueue*, void*); + using TAsyncReaderWriter = grpc::ClientAsyncReaderWriter; + using TAsyncReaderWriterPtr = std::unique_ptr>; + using TAsyncRequest = std::unique_ptr (TStub::*)(grpc::ClientContext*, grpc::CompletionQueue*, void*); explicit TStreamRequestReadWriteProcessor(TConnectedCallback&& callback) : ConnectedCallback(std::move(callback)) @@ -957,7 +974,11 @@ class TStreamRequestReadWriteProcessor { std::unique_lock guard(Mutex); - if (Cancelled || ReadFinished || WriteFinished) { + if (Cancelled) { + status = TGrpcStatus(grpc::StatusCode::CANCELLED, "Write request dropped"); + } else if (HalfCloseRequested) { + status = TGrpcStatus(grpc::StatusCode::FAILED_PRECONDITION, "Client write side is already half-closed"); + } else if (ReadFinished || WriteFinished) { status = TGrpcStatus(grpc::StatusCode::CANCELLED, "Write request dropped"); } else if (WriteActive) { auto& item = WriteQueue.emplace_back(); @@ -977,6 +998,42 @@ class TStreamRequestReadWriteProcessor } } + void WritesDone(TWriteCallback callback) override { + TGrpcStatus status; + bool startWritesDone = false; + + { + std::unique_lock guard(Mutex); + if (Cancelled) { + status = TGrpcStatus(grpc::StatusCode::CANCELLED, "WritesDone dropped"); + } else if (HalfCloseRequested) { + status = TGrpcStatus(grpc::StatusCode::FAILED_PRECONDITION, "Client write side is already half-closed"); + } else if (WriteFinished) { + status = TGrpcStatus(grpc::StatusCode::CANCELLED, "WritesDone dropped"); + } else if (WriteActive) { + HalfCloseRequested = true; + auto& item = WriteQueue.emplace_back(); + item.Callback.swap(callback); + item.IsWritesDone = true; + } else { + HalfCloseRequested = true; + WriteActive = true; + WriteDonePending = true; + WriteCallback.swap(callback); + startWritesDone = true; + } + } + + if (startWritesDone) { + Stream->WritesDone(OnWriteDoneTag.Prepare()); + } + if (!status.Ok() && callback) { + RunGuarded([&] { + callback(std::move(status)); + }); + } + } + void ReadInitialMetadata(std::unordered_multimap* metadata, TReadCallback callback) override { TGrpcStatus status; @@ -1098,6 +1155,7 @@ class TStreamRequestReadWriteProcessor private: template friend class TServiceConnection; + friend struct TStreamRequestReadWriteProcessorTestAccess; void Start(TStub& stub, TAsyncRequest asyncRequest, IQueueClientContextProvider* provider) { InitCallbackGuard(provider); @@ -1195,12 +1253,16 @@ class TStreamRequestReadWriteProcessor Y_ABORT_UNLESS(WriteActive, "Unexpected Write done callback"); Y_ABORT_UNLESS(!WriteFinished, "Unexpected WriteFinished flag"); + const bool wasWritesDone = WriteDonePending; + WriteDonePending = false; + if (ok) { okCallback.swap(WriteCallback); } else if (WriteCallback) { // Put callback back on the queue until OnFinished auto& item = WriteQueue.emplace_front(); item.Callback.swap(WriteCallback); + item.IsWritesDone = wasWritesDone; } if (!ok || Cancelled) { @@ -1209,9 +1271,22 @@ class TStreamRequestReadWriteProcessor if (ReadFinished) { Stream->Finish(&Status, OnFinishedTag.Prepare()); } + } else if (wasWritesDone) { + // Client write side is half-closed; no further Write/WritesDone. + WriteActive = false; + WriteFinished = true; + if (ReadFinished) { + Stream->Finish(&Status, OnFinishedTag.Prepare()); + } } else if (!WriteQueue.empty()) { - WriteCallback.swap(WriteQueue.front().Callback); - Stream->Write(WriteQueue.front().Request, OnWriteDoneTag.Prepare()); + auto& next = WriteQueue.front(); + WriteCallback.swap(next.Callback); + if (next.IsWritesDone) { + WriteDonePending = true; + Stream->WritesDone(OnWriteDoneTag.Prepare()); + } else { + Stream->Write(next.Request, OnWriteDoneTag.Prepare()); + } WriteQueue.pop_front(); } else { WriteActive = false; @@ -1319,6 +1394,7 @@ class TStreamRequestReadWriteProcessor struct TWriteItem { TWriteCallback Callback; TRequest Request; + bool IsWritesDone = false; }; private: @@ -1345,6 +1421,8 @@ class TStreamRequestReadWriteProcessor bool ReadFinished = false; bool WriteActive = false; bool WriteFinished = false; + bool HalfCloseRequested = false; + bool WriteDonePending = false; bool Finished = false; bool Cancelled = false; bool FinishedOk = false; diff --git a/src/library/kafka/kafka_records.cpp b/src/library/kafka/kafka_records.cpp index 6dc10a21258..156f5ce1031 100644 --- a/src/library/kafka/kafka_records.cpp +++ b/src/library/kafka/kafka_records.cpp @@ -3,6 +3,7 @@ #include #include +#include #include #include @@ -734,6 +735,12 @@ i32 TKafkaRecordBatch::Size(TKafkaVersion _version) const { return _collector.Size; } +i64 GetRecordTimestamp(i64 baseTimestamp, i64 timestampDelta) { + // Kafka uses Java long addition, which wraps modulo 2^64. Add unsigned + // values to avoid C++ signed overflow, then interpret the resulting bits. + return std::bit_cast(static_cast(baseTimestamp) + static_cast(timestampDelta)); +} + // // TKafkaBatchHeader diff --git a/src/library/kafka/kafka_records.h b/src/library/kafka/kafka_records.h index e0d45ec751a..529e33090ed 100644 --- a/src/library/kafka/kafka_records.h +++ b/src/library/kafka/kafka_records.h @@ -4,6 +4,7 @@ #include #include +#include namespace NKafka { @@ -686,6 +687,8 @@ TString WriteKafkaRecordBatch(const TKafkaRecordBatch& batch, TKafkaVersion vers std::pair GetBatchBaseSeqNo(const TKafkaBatchHeader& header); std::pair GetBatchMaxSeqNo(const TKafkaBatchHeader& header, ui64 baseSeqNo); +// Reconstruct a record timestamp with Kafka's wrapping 64-bit arithmetic. +i64 GetRecordTimestamp(i64 baseTimestamp, i64 timestampDelta); ui64 GetRecordSeqNo(const TKafkaRecordBatch& batch, size_t recordIndex, const TKafkaRecord& record); } // namespace NKafka diff --git a/src/library/kafka/ut/kafka_records_ut.cpp b/src/library/kafka/ut/kafka_records_ut.cpp index 5ecec452593..174177392fa 100644 --- a/src/library/kafka/ut/kafka_records_ut.cpp +++ b/src/library/kafka/ut/kafka_records_ut.cpp @@ -4,6 +4,9 @@ #include +#include +#include + namespace NKafka { namespace { @@ -468,6 +471,49 @@ Y_UNIT_TEST_SUITE(KafkaRecords) { AssertRecordBatchRoundTrip(ECompressionType::ZSTD); } + Y_UNIT_TEST(RecordBatchTimestampsWrapLikeKafka) { + constexpr auto minTimestamp = std::numeric_limits::min(); + constexpr auto maxTimestamp = std::numeric_limits::max(); + for (const auto compressionType : {ECompressionType::NONE, ECompressionType::GZIP, ECompressionType::ZSTD}) { + for (const auto [baseTimestamp, timestampDelta, expectedTimestamp] : { + std::tuple{maxTimestamp, 1, minTimestamp}, + {minTimestamp, -1, maxTimestamp}, + {1, maxTimestamp, minTimestamp}, + {-1, minTimestamp, maxTimestamp}, + {maxTimestamp, maxTimestamp, -2}, + {minTimestamp, minTimestamp, 0}, + {maxTimestamp, 0, maxTimestamp}, + {minTimestamp, 0, minTimestamp}, + {maxTimestamp - 1, 1, maxTimestamp}, + {minTimestamp + 1, -1, minTimestamp}, + {-1, maxTimestamp, maxTimestamp - 1}, + {0, minTimestamp, minTimestamp}, + {maxTimestamp, minTimestamp, -1}, + {minTimestamp, maxTimestamp, -1}, + {1000, 25, 1025}, + {1000, -25, 975}, + {-1, 0, -1}}) + { + auto batch = MakeRecordBatch(compressionType); + batch.BaseTimestamp = baseTimestamp; + batch.MaxTimestamp = expectedTimestamp; + batch.Records = {MakeRecord(timestampDelta, 0, "key", "value")}; + const auto serialized = WriteKafkaRecordBatch(batch); + const auto header = ReadKafkaBatchHeader(serialized); + UNIT_ASSERT(header); + UNIT_ASSERT_VALUES_EQUAL(header->BaseTimestamp, baseTimestamp); + const auto parsed = ReadKafkaRecordBatch(serialized); + + UNIT_ASSERT_VALUES_EQUAL(parsed.BaseTimestamp, baseTimestamp); + UNIT_ASSERT_VALUES_EQUAL(parsed.Records.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(parsed.Records.front().TimestampDelta, timestampDelta); + UNIT_ASSERT_VALUES_EQUAL( + GetRecordTimestamp(parsed.BaseTimestamp, parsed.Records.front().TimestampDelta), expectedTimestamp); + UNIT_ASSERT(KafkaBytesEqual(parsed.Records.front().Value, batch.Records.front().Value)); + } + } + } + Y_UNIT_TEST(SetKafkaBatchBaseOffset) { const TKafkaRecordBatch expected = MakeRecordBatch(ECompressionType::ZSTD); TString batchBytes = WriteKafkaRecordBatch(expected); diff --git a/src/version.h b/src/version.h index 7c41db36087..adff546971d 100644 --- a/src/version.h +++ b/src/version.h @@ -2,7 +2,7 @@ namespace NYdb { -inline const char* YDB_SDK_VERSION = "3.23.0"; +inline const char* YDB_SDK_VERSION = "3.24.0"; inline const char* YDB_CERTIFICATE_FILE_KEY = "ydb_root_ca_v3.pem"; } // namespace NYdb diff --git a/tests/integration/CMakeLists.txt b/tests/integration/CMakeLists.txt index 3bcef103686..93ae62b68f2 100644 --- a/tests/integration/CMakeLists.txt +++ b/tests/integration/CMakeLists.txt @@ -8,3 +8,4 @@ add_subdirectory(server_restart) add_subdirectory(sessions) add_subdirectory(sessions_pool) add_subdirectory(topic) +add_subdirectory(embedding) diff --git a/tests/integration/embedding/CMakeLists.txt b/tests/integration/embedding/CMakeLists.txt new file mode 100644 index 00000000000..b8e9cb376b1 --- /dev/null +++ b/tests/integration/embedding/CMakeLists.txt @@ -0,0 +1,8 @@ +add_ydb_test(NAME client-embedding_it GTEST + SOURCES embedding_it.cpp + LINK_LIBRARIES + YDB-CPP-SDK::Query + YDB-CPP-SDK::Params + YDB-CPP-SDK::Value + LABELS integration +) diff --git a/tests/integration/embedding/embedding_it.cpp b/tests/integration/embedding/embedding_it.cpp new file mode 100644 index 00000000000..bfeb42f9265 --- /dev/null +++ b/tests/integration/embedding/embedding_it.cpp @@ -0,0 +1,65 @@ +#include +#include +#include + +#include + +#include +#include +#include +#include +#include +#include + +namespace NYdb { +namespace { + +void CheckEmbedding(NQuery::TQueryClient& client, const std::string& type, + const TValue& values, const TValue& embedding) { + SCOPED_TRACE(type); + + const auto query = std::format(R"( + DECLARE $values AS List<{}>; + DECLARE $embedding AS Bytes; + SELECT $embedding = Untag(Knn::ToBinaryStringFloat( + ListMap($values, ($value) -> (CAST($value AS Float))) + ), "FloatVector"); + )", type); + const auto params = TParamsBuilder() + .AddParam("$values", values) + .AddParam("$embedding", embedding) + .Build(); + + auto result = client.ExecuteQuery(query, NQuery::TTxControl::NoTx(), params).ExtractValueSync(); + ASSERT_TRUE(result.IsSuccess()) << result.GetIssues().ToString(); + auto parser = result.GetResultSetParser(0); + ASSERT_TRUE(parser.TryNextRow()); + EXPECT_TRUE(parser.ColumnParser(0).GetBool()); +} + +} // namespace + +TEST(Embedding, MatchesKnnSerialization) { + TDriver driver(TDriverConfig() + .SetEndpoint(std::getenv("YDB_ENDPOINT")) + .SetDatabase(std::getenv("YDB_DATABASE"))); + NQuery::TQueryClient client(driver); + + CheckEmbedding(client, "Float", + TValueBuilder().EmptyList(TTypeBuilder().Primitive(EPrimitiveType::Float).Build()).Build(), + NValueHelpers::Embedding(std::vector{})); + CheckEmbedding(client, "Int64", + TValueBuilder().BeginList().AddListItem().Int64(-2).AddListItem().Int64(16777217).EndList().Build(), + NValueHelpers::Embedding(std::array{-2, 16777217})); + CheckEmbedding(client, "Uint16", + TValueBuilder().BeginList().AddListItem().Uint16(2).EndList().Build(), + NValueHelpers::Embedding(std::array{2})); + CheckEmbedding(client, "Float", + TValueBuilder().BeginList().AddListItem().Float(-2.5f).AddListItem().Float(1.25f).EndList().Build(), + NValueHelpers::Embedding(std::vector{-2.5f, 1.25f})); + CheckEmbedding(client, "Double", + TValueBuilder().BeginList().AddListItem().Double(-2.5).AddListItem().Double(1.00000001).EndList().Build(), + NValueHelpers::Embedding(std::vector{-2.5, 1.00000001})); +} + +} // namespace NYdb diff --git a/tests/unit/client/CMakeLists.txt b/tests/unit/client/CMakeLists.txt index 19ab2f4af09..1874d8a6b77 100644 --- a/tests/unit/client/CMakeLists.txt +++ b/tests/unit/client/CMakeLists.txt @@ -2,6 +2,26 @@ add_subdirectory(federated_topic) add_subdirectory(oauth2_token_exchange/helpers) add_subdirectory(topic) +add_ydb_test(NAME client-embedding_ut GTEST + SOURCES value/embedding_ut.cpp + LINK_LIBRARIES YDB-CPP-SDK::Value + LABELS unit +) + +add_ydb_test(NAME client-oidc_ut + SOURCES + oidc/credentials_ut.cpp + oidc/protocol_ut.cpp + oidc/test_server.cpp + oidc/helpers/test_server.cpp + LINK_LIBRARIES + client-types-credentials-oidc + http-server + json + string_utils-base64 + LABELS unit +) + add_ydb_test(NAME client-connection_string_ut GTEST SOURCES connection_string/connection_string_ut.cpp diff --git a/tests/unit/client/driver/driver_ut.cpp b/tests/unit/client/driver/driver_ut.cpp index 382027f2bf9..593f8488092 100644 --- a/tests/unit/client/driver/driver_ut.cpp +++ b/tests/unit/client/driver/driver_ut.cpp @@ -1,6 +1,7 @@ #include #include #include +#include #include #include #include @@ -65,6 +66,11 @@ IGfPhGBVwOMnr+uhwtpj4PAOIrlOQD/fBsaRtYuBRdg2 { BuildInfo = ReadBuildInfo(context); + const auto& metadata = context->client_metadata(); + if (const auto it = metadata.find(YDB_AUTH_TICKET_HEADER); it != metadata.end()) { + AuthTicket.assign(it->second.data(), it->second.length()); + } + std::cerr << "ListEndpoints: " << request->ShortDebugString() << std::endl; const auto* result = MapFindPtr(MockResults, request->database()); @@ -80,6 +86,7 @@ IGfPhGBVwOMnr+uhwtpj4PAOIrlOQD/fBsaRtYuBRdg2 // From database name to result std::unordered_map MockResults; std::string BuildInfo; + std::string AuthTicket; }; class TMockTableService : public Ydb::Table::V1::TableService::Service { @@ -381,6 +388,25 @@ Y_UNIT_TEST_SUITE(DeferredCredentialsTest) { } Y_UNIT_TEST_SUITE(CppGrpcClientSimpleTest) { + Y_UNIT_TEST(OidcTokenIsSentAsBearerTicket) { + TPortManager pm; + TMockDiscoveryService discoveryService; + discoveryService.MockResults["/Root/My/DB"] = {}; + const auto address = TStringBuilder() << "127.0.0.1:" << pm.GetPort(); + auto server = StartGrpcServer(address, discoveryService); + + NOidc::TOidcConfig oidc; + oidc.Issuer = "https://issuer.example"; + oidc.FlowConfig = NOidc::TStaticOidcConfig{.AccessToken = "oidc-access"}; + auto driver = TDriver(TDriverConfig() + .SetEndpoint(address) + .SetDatabase("/Root/My/DB") + .SetDiscoveryMode(EDiscoveryMode::Sync) + .SetCredentialsProviderFactory(NOidc::CreateOidcProviderFactory(oidc))); + + UNIT_ASSERT_VALUES_EQUAL(discoveryService.AuthTicket, "Bearer oidc-access"); + } + Y_UNIT_TEST(ReusesCredentialsProviderForSameIdentity) { std::atomic_int providerCount = 0; auto driver = TDriver( diff --git a/tests/unit/client/oidc/credentials_ut.cpp b/tests/unit/client/oidc/credentials_ut.cpp new file mode 100644 index 00000000000..60ad56575f8 --- /dev/null +++ b/tests/unit/client/oidc/credentials_ut.cpp @@ -0,0 +1,1135 @@ +#include +#include + +#include "test_server.h" +#include +#include +#include +#include + +#include +#include + +#include + +using namespace NYdb; +using namespace NYdb::NOidc; +using NYdb::NOidc::NPrivate::GetOidcClientIdentity; + +namespace { + +class TThrowingOidcFacility: public TQueuedOidcFacility { +public: + void PostToResponseQueue(TPostTaskCb&& callback) override; +}; + +class TFailingOidcCacher: public TMemoryTokenCacher { +public: + TFailingOidcCacher(bool failRead, bool failWrite); + + std::optional Read() const override; + void Write(const TTokenCache& cache) override; + +private: + const bool FailRead; + const bool FailWrite; +}; + +class TGatedOidcAcceptor: public TTestAcceptor { +public: + void Accept(const TDeviceAuthInfo& info) override; + + NThreading::TPromise Release = NThreading::NewPromise(); + NThreading::TPromise Finished = NThreading::NewPromise(); +}; + +void TThrowingOidcFacility::PostToResponseQueue(TPostTaskCb&&) { + throw std::runtime_error("response queue unavailable"); +} + +TFailingOidcCacher::TFailingOidcCacher(bool failRead, bool failWrite) + : FailRead(failRead) + , FailWrite(failWrite) +{ +} + +std::optional TFailingOidcCacher::Read() const { + if (FailRead) { + throw std::runtime_error("cache read failed: private-token"); + } + return TMemoryTokenCacher::Read(); +} + +void TFailingOidcCacher::Write(const TTokenCache& cache) { + if (FailWrite) { + throw std::runtime_error("cache write failed: private-token"); + } + TMemoryTokenCacher::Write(cache); +} + +void TGatedOidcAcceptor::Accept(const TDeviceAuthInfo& info) { + TTestAcceptor::Accept(info); + Release.GetFuture().Wait(); + Finished.TrySetValue(); +} + +} // namespace + +Y_UNIT_TEST_SUITE(TOidcCredentials) { +Y_UNIT_TEST(FactoryReturnsConcreteProviders) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + auto cache = std::make_shared(); + cache->Write({{"cached", std::nullopt}, std::nullopt}); + config.Cacher(cache); + auto facility = std::make_shared(); + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + UNIT_ASSERT(std::dynamic_pointer_cast(provider) != nullptr); + config.FlowConfig = TClientOidcConfig{"client", "secret", {}}; + provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + UNIT_ASSERT(std::dynamic_pointer_cast(provider) != nullptr); + config.FlowConfig = TDeviceOidcConfig{"client", {}}; + provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + UNIT_ASSERT(std::dynamic_pointer_cast(provider) != nullptr); +} + +Y_UNIT_TEST(BearerTicketDoesNotChangeCachedToken) { + auto cache = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{ + .AccessToken = "opaque-access", + .ExpiresAt = TInstant::Now() + TDuration::Hours(1), + }; + config.Cacher(cache); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + auto pending = provider->GetAuthInfoAsync(); + const bool wasPending = !pending.IsReady(); + cache->Release.TrySetValue(); + + UNIT_ASSERT(wasPending); + UNIT_ASSERT_VALUES_EQUAL(pending.GetValueSync(), "Bearer opaque-access"); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer opaque-access"); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value()); + UNIT_ASSERT(!stored->RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(stored->AccessToken.Token, "opaque-access"); +} + +Y_UNIT_TEST(LiveFacilityDiscardCompletesPendingFuture) { + auto cacher = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + config.Cacher(cacher); + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + UNIT_ASSERT(cacher->Entered.GetFuture().Wait(TDuration::Seconds(1))); + auto pending = provider->GetAuthInfoAsync(); + cacher->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + auto reentered = NThreading::NewPromise(); + pending.Subscribe([facility, reentered](const auto&) mutable { + facility->RunTasks(); + reentered.TrySetValue(); + }); + facility->DiscardTasks(); + UNIT_ASSERT(reentered.GetFuture().Wait(TDuration::Seconds(1))); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(1))); + UNIT_ASSERT_EXCEPTION(pending.GetValueSync(), std::exception); +} + +Y_UNIT_TEST(ThrowingDiscardSubscriberDoesNotTerminateWorker) { + auto cacher = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + config.Cacher(cacher); + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + UNIT_ASSERT(cacher->Entered.GetFuture().Wait(TDuration::Seconds(1))); + auto pending = provider->GetAuthInfoAsync(); + cacher->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + auto reentered = NThreading::NewPromise(); + pending.Subscribe([facility, reentered](const auto&) mutable { + facility->RunTasks(); + reentered.TrySetValue(); + throw std::runtime_error("subscriber failure"); + }); + facility->DiscardTasks(); + UNIT_ASSERT(reentered.GetFuture().Wait(TDuration::Seconds(1))); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(1))); + UNIT_ASSERT_EXCEPTION(pending.GetValueSync(), std::exception); +} + +Y_UNIT_TEST(ThrowingSuccessSubscriberDoesNotInterruptResponseQueue) { + auto cache = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + config.Cacher(cache); + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + auto pending = provider->GetAuthInfoAsync(); + pending.Subscribe([](const auto&) { + throw std::runtime_error("subscriber failure"); + }); + cache->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + bool nextTaskRan = false; + facility->PostToResponseQueue([&nextTaskRan] { nextTaskRan = true; }); + UNIT_ASSERT_NO_EXCEPTION(facility->RunTasks()); + UNIT_ASSERT(nextTaskRan); + UNIT_ASSERT_VALUES_EQUAL(pending.GetValueSync(), "Bearer opaque"); +} + +Y_UNIT_TEST(CacheFailuresDoNotPreventClientAuthentication) { + for (const bool failRead : {false, true}) { + for (const bool failWrite : {false, true}) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto cache = std::make_shared(failRead, failWrite); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer access"); + UNIT_ASSERT(provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); + if (!failRead && !failWrite) { + UNIT_ASSERT_VALUES_EQUAL(cache->Read()->AccessToken.Token, "access"); + } + } + } +} + +Y_UNIT_TEST(CacheWriteFailureDoesNotPreventStaticAuthentication) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + config.Cacher(std::make_shared(true, true)); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer opaque"); +} + +Y_UNIT_TEST(CacheFailuresDoNotDiscardDeviceAuthorization) { + TOidcTestServer server; + server.Enqueue(TString("{\"device_code\":\"private-device\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + + server.Issuer() + "/verify\",\"expires_in\":60,\"interval\":1}", HTTP_OK); + server.Enqueue(R"({"access_token":"user-access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto acceptor = std::make_shared(); + auto config = server.ClientConfig().Acceptor(acceptor).Cacher(std::make_shared(true, true)); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer user-access"); + UNIT_ASSERT_VALUES_EQUAL(acceptor->Wait().UserCode, "ABCD"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 2); +} + +Y_UNIT_TEST(MalformedJwtExpiryDoesNotRejectAccessToken) { + const auto token = "e30." + std::string(Base64EncodeUrl(R"({"exp":"unknown"})")) + ".signature"; + for (const bool useStatic : {false, true}) { + TOidcTestServer server; + auto config = server.ClientConfig(); + if (useStatic) { + config.FlowConfig = TStaticOidcConfig{.AccessToken = token}; + } else { + server.Enqueue(TString("{\"access_token\":\"") + token + "\",\"token_type\":\"Bearer\"}", HTTP_OK); + } + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer " + token); + UNIT_ASSERT(provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), useStatic ? 0 : 1); + } +} + +Y_UNIT_TEST(CustomHooksIsolateFactoryIdentity) { + for (const TFlowConfig& flow : { + TFlowConfig{TStaticOidcConfig{.AccessToken = "opaque"}}, + TFlowConfig{TClientOidcConfig{"client", "secret", {}}}}) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = flow; + const auto identity = GetOidcClientIdentity(config); + UNIT_ASSERT_VALUES_EQUAL(CreateOidcProviderFactory(config)->GetClientIdentity(), + CreateOidcProviderFactory(config)->GetClientIdentity()); + for (const bool useCacher : {false, true}) { + auto first = config; + auto second = config; + if (useCacher) { + first.Cacher(std::make_shared()); + second.Cacher(std::make_shared()); + } else { + first.Acceptor(std::make_shared()); + second.Acceptor(std::make_shared()); + } + const auto firstFactory = CreateOidcProviderFactory(first); + const auto firstIdentity = firstFactory->GetClientIdentity(); + UNIT_ASSERT(firstIdentity != CreateOidcProviderFactory(second)->GetClientIdentity()); + UNIT_ASSERT_VALUES_EQUAL(firstIdentity, firstFactory->GetClientIdentity()); + UNIT_ASSERT_VALUES_EQUAL(GetOidcClientIdentity(first), identity); + } + } +} + +Y_UNIT_TEST(ClientRetriesTransientError) { + TOidcTestServer server; + auto replyGate = NThreading::NewPromise(); + server.BlockTokenRepliesUntil(replyGate.GetFuture()); + server.Enqueue(R"({"error":"temporarily_unavailable"})", HTTP_SERVICE_UNAVAILABLE); + server.Enqueue(R"({"access_token":"retried","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(server.ClientConfig())->CreateProvider(facility); + auto token = provider->GetAuthInfoAsync(); + replyGate.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + facility->RunTasks(); + UNIT_ASSERT(token.Wait(TDuration::Seconds(5))); + UNIT_ASSERT_EXCEPTION_CONTAINS(token.GetValueSync(), std::exception, "503"); + if (!provider->GetAuthInfoAsync().HasValue()) { + UNIT_ASSERT(facility->WaitForTask()); + facility->RunTasks(); + } + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer retried"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 2); +} + +Y_UNIT_TEST(JwtExpiryIsOnlyASchedulingHint) { + using NYdb::NOidc::NPrivate::JwtExpiry; + const std::string token = "e30." + std::string(Base64EncodeUrl(R"({"exp":2000000000})")) + ".signature"; + UNIT_ASSERT(JwtExpiry(token).has_value()); + UNIT_ASSERT_VALUES_EQUAL(*JwtExpiry(token), TInstant::Seconds(2000000000)); + UNIT_ASSERT(!JwtExpiry("opaque").has_value()); + UNIT_ASSERT(!JwtExpiry("e30.invalid.signature").has_value()); +} + +Y_UNIT_TEST(IndependentStaticFactories) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + TStaticOidcConfig first; + first.AccessToken = "first-opaque-token"; + config.FlowConfig = first; + auto firstFactory = CreateOidcProviderFactory(config); + first.AccessToken = "second-opaque-token"; + config.FlowConfig = first; + auto secondFactory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(firstFactory->CreateProvider()->GetAuthInfo(), "Bearer first-opaque-token"); + UNIT_ASSERT_VALUES_EQUAL(secondFactory->CreateProvider()->GetAuthInfo(), "Bearer second-opaque-token"); + UNIT_ASSERT(firstFactory->GetClientIdentity() != secondFactory->GetClientIdentity()); +} + +Y_UNIT_TEST(ClientIdentityDoesNotExposeCredentials) { + const std::string clientSecret = "private-client-secret"; + const std::string staticToken = "private-static-access-token"; + const std::string cachedAccessToken = "private-cached-access-token"; + const std::string refreshToken = "private-cached-refresh-token"; + auto cache = std::make_shared(); + cache->Write({{cachedAccessToken, std::nullopt}, TOAuthToken{refreshToken, std::nullopt}}); + + for (const TFlowConfig& flow : { + TFlowConfig{TStaticOidcConfig{staticToken, std::nullopt}}, + TFlowConfig{TClientOidcConfig{"client", clientSecret, {"read"}}}, + TFlowConfig{TDeviceOidcConfig{"client", {"read"}}}}) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = flow; + config.Cacher(cache); + const auto factory = CreateOidcProviderFactory(config); + for (const auto& identity : {GetOidcClientIdentity(config), factory->GetClientIdentity()}) { + for (const auto& secret : {clientSecret, staticToken, cachedAccessToken, refreshToken}) { + UNIT_ASSERT(identity.find(secret) == std::string::npos); + } + } + } +} + +Y_UNIT_TEST(RejectsMissingClientSecret) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + TClientOidcConfig flow; + flow.ClientId = "client"; + config.FlowConfig = flow; + UNIT_ASSERT_EXCEPTION(CreateOidcProviderFactory(config), std::invalid_argument); +} + +Y_UNIT_TEST(RejectsInsecureIssuer) { + TOidcTestServer server; + auto config = server.ClientConfig(); + config.Issuer = "http://issuer.example"; + UNIT_ASSERT_EXCEPTION(CreateOidcProviderFactory(config), std::invalid_argument); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); +} + +Y_UNIT_TEST(UrlComponentsPreservePortAndQuery) { + using NOidc::NPrivate::ParseUrl; + const auto endpoint = ParseUrl("https://issuer.example:8443/token?scope=a%2Bb&x=1", false); + UNIT_ASSERT_VALUES_EQUAL(endpoint.GetPort(), 8443); + UNIT_ASSERT_VALUES_EQUAL(endpoint.PrintS(NUri::TUri::FlagScheme | NUri::TUri::FlagHost | NUri::TUri::FlagHostAscii), "https://issuer.example"); + UNIT_ASSERT_VALUES_EQUAL(endpoint.PrintS(NUri::TUri::FlagPath | NUri::TUri::FlagQuery), "/token?scope=a%2Bb&x=1"); + const auto root = ParseUrl("https://issuer.example", true); + UNIT_ASSERT_VALUES_EQUAL(root.GetPort(), 443); + UNIT_ASSERT_VALUES_EQUAL(root.PrintS(NUri::TUri::FlagPath | NUri::TUri::FlagQuery), "/"); + const auto query = ParseUrl("https://issuer.example?code=abc", false); + UNIT_ASSERT_VALUES_EQUAL(query.PrintS(NUri::TUri::FlagPath | NUri::TUri::FlagQuery), "/?code=abc"); +} + +Y_UNIT_TEST(UrlValidationRejectsInvalidOidcEndpoints) { + using NOidc::NPrivate::ParseUrl; + for (const char* url : { + "", "http://issuer.example", "https:///", "https://user:pass@issuer.example/token", + "https://issuer.example/token#fragment", "https://issuer.example/a b", + "https://issuer.example:65536/token"}) { + UNIT_ASSERT_EXCEPTION(ParseUrl(url, false), std::invalid_argument); + } + UNIT_ASSERT_EXCEPTION(ParseUrl("https://issuer.example?query=value", true), std::invalid_argument); +} + +Y_UNIT_TEST(ClientGrantEncodesForm) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto factory = CreateOidcProviderFactory(server.ClientConfig()); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer access"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Method, "POST"); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("grant_type"), "client_credentials"); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Authorization, "Basic " + Base64Encode("client:secret+%2B%26")); + UNIT_ASSERT(!requests[0].Form.Has("client_secret")); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("scope"), "read write openid"); +} + +Y_UNIT_TEST(ClientGrantIncludesOpenidOnce) { + const std::vector, TString>> cases = { + {{}, "openid"}, + {{"openid"}, "openid"}, + {{"profile", "openid"}, "profile openid"}, + }; + for (const auto& [scopes, expected] : cases) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto config = server.ClientConfig(); + std::get(config.FlowConfig).Scopes = scopes; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer access"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests()[0].Form.Get("scope"), expected); + } +} + +Y_UNIT_TEST(ClientSecretBasic) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"bEaReR","expires_in":600})", HTTP_OK); + auto config = server.ClientConfig(); + auto& flow = std::get(config.FlowConfig); + flow.ClientId = "client:+ %"; + flow.ClientSecret = "secret:/?#[]@!$&'()*+,;= %~_-.Я"; + auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer access"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Authorization, "Basic " + Base64Encode( + "client%3A%2B+%25:secret%3A%2F%3F%23%5B%5D%40%21%24%26%27%28%29%2A%2B%2C%3B%3D+%25~_-.%D0%AF")); + UNIT_ASSERT(!requests[0].Form.Has("client_secret")); +} + +Y_UNIT_TEST(ReusesCachedAccessWithoutDiscovery) { + TOidcTestServer server; + auto cache = std::make_shared(); + cache->Write({{"cached", TInstant::Now() + TDuration::Hours(1)}, std::nullopt}); + auto config = server.ClientConfig().Cacher(cache); + auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer cached"); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); +} + +Y_UNIT_TEST(RefreshesExpiredCacheAndRetainsRefreshToken) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"refresh", std::nullopt}}); + auto config = server.ClientConfig().Cacher(cache); + auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer fresh"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("grant_type"), "refresh_token"); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("refresh_token"), "refresh"); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value()); + UNIT_ASSERT(stored->RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(stored->RefreshToken->Token, "refresh"); +} + +Y_UNIT_TEST(RetainedRefreshTokenReceivesUpdatedExpiry) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600,"refresh_expires_in":1200})", HTTP_OK); + auto cache = std::make_shared(); + const auto now = TInstant::Now(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"refresh", now + TDuration::Seconds(60)}}); + auto factory = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache)); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer fresh"); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value() && stored->RefreshToken.has_value() && stored->RefreshToken->ExpiresAt.has_value()); + UNIT_ASSERT(*stored->RefreshToken->ExpiresAt >= now + TDuration::Seconds(1200)); +} + +Y_UNIT_TEST(InvalidRefreshFallsBackToClientGrant) { + TOidcTestServer server; + server.Enqueue(R"({"error":"invalid_grant"})", HTTP_BAD_REQUEST); + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600,"refresh_token":"rotated"})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"refresh", std::nullopt}}); + auto config = server.ClientConfig().Cacher(cache); + auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer fresh"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 2); + UNIT_ASSERT_VALUES_EQUAL(requests[1].Form.Get("grant_type"), "client_credentials"); + UNIT_ASSERT_VALUES_EQUAL(cache->Read()->RefreshToken->Token, "rotated"); +} + +Y_UNIT_TEST(ExpiredStaticTokenDoesNotContactIssuer) { + TOidcTestServer server; + auto config = server.ClientConfig(); + config.FlowConfig = TStaticOidcConfig{"expired", TInstant::Seconds(1)}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + auto result = provider->GetAuthInfoAsync(); + UNIT_ASSERT(result.Wait(TDuration::Seconds(1))); + UNIT_ASSERT(result.HasException()); + UNIT_ASSERT(!provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(RejectsMalformedResponsesWithoutLeakingSecrets) { + for (const TString& body : { + TString("secret-response-is-not-json"), + TString(R"({"access_token":"secret-access","token_type":"unsupported-secret","expires_in":600})"), + TString(R"({"access_token":"secret-access","token_type":"Bearer","expires_in":0})")}) { + TOidcTestServer server; + server.Enqueue(body, HTTP_OK); + auto factory = CreateOidcProviderFactory(server.ClientConfig()); + try { + factory->CreateProvider()->GetAuthInfo(); + UNIT_FAIL("Expected invalid token response"); + } catch (const std::exception& error) { + UNIT_ASSERT(!TString(error.what()).Contains("secret")); + } + } +} + +Y_UNIT_TEST(ConcurrentRequestsShareInitialGrant) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto factory = CreateOidcProviderFactory(server.ClientConfig()); + auto provider = factory->CreateProvider(); + std::vector> requests; + for (size_t i = 0; i < 8; ++i) { + requests.push_back(std::async(std::launch::async, [provider] { return provider->GetAuthInfo(); })); + } + for (auto& request : requests) { + UNIT_ASSERT_VALUES_EQUAL(request.get(), "Bearer access"); + } + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); +} + +Y_UNIT_TEST(DeviceAuthorization) { + TOidcTestServer server; + server.Enqueue(TString("{\"device_code\":\"private-device\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + server.Issuer() + "/verify\",\"verification_uri_complete\":\"" + server.Issuer() + "/verify?user_code=ABCD\",\"expires_in\":60,\"interval\":1}", HTTP_OK); + server.Enqueue(R"({"access_token":"user-access","token_type":"Bearer","expires_in":600,"refresh_token":"user-refresh"})", HTTP_OK); + auto acceptor = std::make_shared(); + auto cache = std::make_shared(); + auto config = server.ClientConfig().Acceptor(acceptor).Cacher(cache); + config.FlowConfig = TDeviceOidcConfig{"public-client", {"openid"}}; + auto factory = CreateOidcProviderFactory(config); + auto provider = factory->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer user-access"); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value() && stored->RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(stored->AccessToken.Token, "user-access"); + UNIT_ASSERT_VALUES_EQUAL(stored->RefreshToken->Token, "user-refresh"); + const auto info = acceptor->Wait(); + UNIT_ASSERT_VALUES_EQUAL(info.UserCode, "ABCD"); + UNIT_ASSERT_VALUES_EQUAL(info.VerificationUrl, server.Issuer() + "/verify"); + UNIT_ASSERT(info.VerificationUrlComplete.has_value()); + UNIT_ASSERT_VALUES_EQUAL(info.VerificationUrlComplete.value(), server.Issuer() + "/verify?user_code=ABCD"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 2); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("scope"), "openid"); + UNIT_ASSERT_VALUES_EQUAL(requests[1].Form.Get("grant_type"), "urn:ietf:params:oauth:grant-type:device_code"); + UNIT_ASSERT_VALUES_EQUAL(requests[1].Form.Get("device_code"), "private-device"); +} + +Y_UNIT_TEST(StoppingDeviceWaitCompletesPendingFuture) { + TOidcTestServer server; + server.Enqueue(TString("{\"device_code\":\"private-device\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + server.Issuer() + "/verify\",\"expires_in\":600,\"interval\":300}", HTTP_OK); + auto acceptor = std::make_shared(); + auto config = server.ClientConfig().Acceptor(acceptor); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + auto factory = CreateOidcProviderFactory(config); + auto facility = CreateSimpleCoreFacility(); + auto provider = factory->CreateProvider(facility); + const auto pending = provider->GetAuthInfoAsync(); + UNIT_ASSERT(!acceptor->Wait().VerificationUrlComplete.has_value()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests()[0].Form.Get("scope"), "openid"); + provider.reset(); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(2))); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); +} + +Y_UNIT_TEST(DeviceWithoutAcceptorCanUseCache) { + TOidcTestServer server; + auto cache = std::make_shared(); + cache->Write({{"cached-user", TInstant::Now() + TDuration::Hours(1)}, std::nullopt}); + auto config = server.ClientConfig().Cacher(cache); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT_VALUES_EQUAL(factory->CreateProvider()->GetAuthInfo(), "Bearer cached-user"); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); +} + +Y_UNIT_TEST(ExpiredFacilityCompletesPendingFuture) { + TOidcTestServer server; + auto factory = CreateOidcProviderFactory(server.ClientConfig()); + auto provider = factory->CreateProvider(std::weak_ptr{}); + auto pending = provider->GetAuthInfoAsync(); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(1))); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); +} + +Y_UNIT_TEST(StaticTokenExpiryNeverLeavesPendingRequest) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + TStaticOidcConfig flow; + flow.AccessToken = "static"; + flow.ExpiresAt = TInstant::Now() + TDuration::MilliSeconds(100); + config.FlowConfig = flow; + auto factory = CreateOidcProviderFactory(config); + auto provider = factory->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer static"); + // A real expiration boundary is the behavior under test. Wait for that + // known deadline rather than estimating background worker progress. + NThreading::NewPromise().GetFuture().Wait(*flow.ExpiresAt + TDuration::MilliSeconds(1)); + auto expired = provider->GetAuthInfoAsync(); + UNIT_ASSERT(expired.Wait(TDuration::Seconds(1))); + UNIT_ASSERT(expired.HasException()); +} + +Y_UNIT_TEST(DevicePollingPolicy) { + using namespace NOidc::NPrivate; + TDevicePolling polling; + const auto now = TInstant::Seconds(1'000); + polling.Deadline = now + TDuration::Seconds(20); + UNIT_ASSERT_VALUES_EQUAL(polling.NextDelay(now), TDuration::Seconds(5)); + polling.HandleError(TError("pending", false, "authorization_pending")); + UNIT_ASSERT_VALUES_EQUAL(polling.NextDelay(now), TDuration::Seconds(5)); + polling.HandleError(TError("slow", false, "slow_down")); + UNIT_ASSERT_VALUES_EQUAL(polling.NextDelay(now), TDuration::Seconds(10)); + polling.HandleError(TError("transport", true, {})); + UNIT_ASSERT_VALUES_EQUAL(polling.NextDelay(now), TDuration::Seconds(20)); + UNIT_ASSERT_VALUES_EQUAL(polling.NextDelay(now + TDuration::Seconds(19)), TDuration::Seconds(1)); + UNIT_ASSERT_EXCEPTION(polling.NextDelay(polling.Deadline), TError); + UNIT_ASSERT_EXCEPTION(polling.HandleError(TError("denied", false, "access_denied")), TError); +} + +Y_UNIT_TEST(EquivalentFactoriesHaveStableIdentity) { + for (const TFlowConfig& flow : { + TFlowConfig{TStaticOidcConfig{.AccessToken = "opaque"}}, + TFlowConfig{TClientOidcConfig{"client", "secret", {"read"}}}, + TFlowConfig{TDeviceOidcConfig{"client", {"read"}}}}) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = flow; + for (const bool useHooks : {false, true}) { + if (useHooks) { + config.Cacher(std::make_shared()); + config.Acceptor(std::make_shared()); + } + UNIT_ASSERT_VALUES_EQUAL(CreateOidcProviderFactory(config)->GetClientIdentity(), + CreateOidcProviderFactory(config)->GetClientIdentity()); + } + } +} + +Y_UNIT_TEST(DeviceFactoriesDoNotShareUserIdentity) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TDeviceOidcConfig{"public-client", {"openid"}}; + auto alice = std::make_shared(); + auto bob = std::make_shared(); + auto aliceFactory = CreateOidcProviderFactory(config.Cacher(alice)); + auto bobFactory = CreateOidcProviderFactory(config.Cacher(bob)); + UNIT_ASSERT(aliceFactory->GetClientIdentity() != bobFactory->GetClientIdentity()); + UNIT_ASSERT_VALUES_EQUAL(aliceFactory->GetClientIdentity(), aliceFactory->GetClientIdentity()); +} + +Y_UNIT_TEST(StaticTokenReplacesCachedCredentials) { + TOidcTestServer server; + auto cache = std::make_shared(); + cache->Write({{"cached", TInstant::Now() + TDuration::Hours(1)}, TOAuthToken{"refresh", std::nullopt}}); + auto config = server.ClientConfig().Cacher(cache); + const auto expiry = TInstant::Now() + TDuration::Seconds(1); + config.FlowConfig = TStaticOidcConfig{"initial", expiry}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer initial"); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value()); + UNIT_ASSERT_VALUES_EQUAL(stored->AccessToken.Token, "initial"); + UNIT_ASSERT(stored->AccessToken.ExpiresAt.has_value()); + UNIT_ASSERT_VALUES_EQUAL(*stored->AccessToken.ExpiresAt, expiry); + UNIT_ASSERT(!stored->RefreshToken.has_value()); + NThreading::NewPromise().GetFuture().Wait(expiry + TDuration::MilliSeconds(1)); + auto result = provider->GetAuthInfoAsync(); + UNIT_ASSERT(result.Wait(TDuration::Seconds(1))); + UNIT_ASSERT(result.HasException()); + UNIT_ASSERT(!provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(ClientRefreshPersistsTokensForNextProvider) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600,"refresh_token":"rotated"})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"initial-refresh", std::nullopt}}); + const auto config = server.ClientConfig().Cacher(cache); + { + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer fresh"); + } + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value() && stored->RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(stored->AccessToken.Token, "fresh"); + UNIT_ASSERT_VALUES_EQUAL(stored->RefreshToken->Token, "rotated"); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer fresh"); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 1); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("grant_type"), "refresh_token"); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("refresh_token"), "initial-refresh"); +} + +Y_UNIT_TEST(StandaloneDestructionWaitsForResponseCallback) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto replyGate = NThreading::NewPromise(); + server.BlockTokenRepliesUntil(replyGate.GetFuture()); + auto factory = CreateOidcProviderFactory(server.ClientConfig()); + auto provider = factory->CreateProvider(); + auto entered = NThreading::NewPromise(); + auto release = NThreading::NewPromise(); + provider->GetAuthInfoAsync().Subscribe([entered, release](const auto&) mutable { + entered.TrySetValue(); + release.GetFuture().Wait(); + }); + replyGate.TrySetValue(); + const bool callbackStarted = entered.GetFuture().Wait(TDuration::Seconds(5)); + auto stopped = std::async(std::launch::async, + [provider = std::move(provider), factory = std::move(factory)]() mutable { + factory.reset(); + provider.reset(); + }); + const bool completed = stopped.wait_for(std::chrono::milliseconds(100)) == std::future_status::ready; + release.TrySetValue(); + stopped.get(); + UNIT_ASSERT(callbackStarted); + UNIT_ASSERT(!completed); +} + +Y_UNIT_TEST(DiscardedCompletionAllowsExternalDestruction) { + auto cache = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + config.Cacher(cache); + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + auto pending = provider->GetAuthInfoAsync(); + auto finished = NThreading::NewPromise(); + pending.Subscribe([finished](const auto&) mutable { + finished.TrySetValue(); + }); + cache->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + facility->DiscardTasks(); + UNIT_ASSERT(finished.GetFuture().Wait(TDuration::Seconds(5))); + provider.reset(); + UNIT_ASSERT(pending.HasException()); +} + +Y_UNIT_TEST(DestructionJoinsWorkerUntilAcceptorReturns) { + TOidcTestServer server; + server.Enqueue(TString("{\"device_code\":\"private-device\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + + server.Issuer() + "/verify\",\"expires_in\":60,\"interval\":1}", HTTP_OK); + auto acceptor = std::make_shared(); + auto config = server.ClientConfig().Acceptor(acceptor); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + auto pending = provider->GetAuthInfoAsync(); + acceptor->Wait(); + auto stopped = std::async(std::launch::async, [provider = std::move(provider)]() mutable { + provider.reset(); + }); + const bool cancelled = pending.Wait(TDuration::Seconds(5)); + const bool completed = stopped.wait_for(std::chrono::milliseconds(100)) == std::future_status::ready; + acceptor->Release.TrySetValue(); + stopped.get(); + UNIT_ASSERT(cancelled); + UNIT_ASSERT(!completed); + UNIT_ASSERT(acceptor->Finished.GetFuture().HasValue()); + UNIT_ASSERT(pending.HasException()); +} + +Y_UNIT_TEST(StandaloneSubscriberCanWaitForAnotherProvider) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{.AccessToken = "opaque"}; + auto firstCache = std::make_shared(); + auto secondCache = std::make_shared(); + auto first = CreateOidcProviderFactory(config.Cacher(firstCache))->CreateProvider(); + auto second = CreateOidcProviderFactory(config.Cacher(secondCache))->CreateProvider(); + auto entered = NThreading::NewPromise(); + auto finished = NThreading::NewPromise(); + first->GetAuthInfoAsync().Subscribe([second, entered, finished](const auto&) mutable { + auto pending = second->GetAuthInfoAsync(); + entered.TrySetValue(); + finished.TrySetValue(pending.Wait(TDuration::Seconds(2))); + }); + firstCache->Release.TrySetValue(); + const bool callbackStarted = entered.GetFuture().Wait(TDuration::Seconds(5)); + secondCache->Release.TrySetValue(); + UNIT_ASSERT(callbackStarted); + UNIT_ASSERT(finished.GetFuture().Wait(TDuration::Seconds(5))); + UNIT_ASSERT(finished.GetFuture().GetValue()); +} + +Y_UNIT_TEST(DeviceRejectsStreamingTokenReceivedAfterExpiry) { + TOidcTestServer server; + server.Enqueue(TString("{\"device_code\":\"code\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + + server.Issuer() + "/verify\",\"expires_in\":2,\"interval\":1}", HTTP_OK); + server.Enqueue(R"({"access_token":"late","token_type":"Bearer","expires_in":600})", HTTP_OK); + server.SetTokenReplyDelay(TDuration::MilliSeconds(100)); + auto config = server.ClientConfig().Acceptor(std::make_shared()); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + auto provider = CreateOidcProviderFactory(config)->CreateProvider(); + auto pending = provider->GetAuthInfoAsync(); + const bool requested = server.WaitRequests(2); + const bool completed = pending.Wait(TDuration::Seconds(10)); + provider.reset(); + UNIT_ASSERT(requested); + UNIT_ASSERT(completed); + UNIT_ASSERT_EXCEPTION_CONTAINS(pending.GetValueSync(), std::exception, "device authorization expired"); +} + +Y_UNIT_TEST(DeviceTokenRequestUsesRemainingLifetime) { + auto gate = NThreading::NewPromise(); + // Request deadline is exercised via the device deadline, which must + // bound a token request even when the socket timeout is much longer. + // The device code expires while the server holds its token response. + TOidcTestServer deviceServer; + deviceServer.Enqueue(TString("{\"device_code\":\"code\",\"user_code\":\"ABCD\",\"verification_uri\":\"") + deviceServer.Issuer() + "/verify\",\"expires_in\":2,\"interval\":1}", HTTP_OK); + deviceServer.Enqueue(R"({"access_token":"late","token_type":"Bearer","expires_in":600})", HTTP_OK); + deviceServer.BlockTokenRepliesUntil(gate.GetFuture()); + auto config = deviceServer.ClientConfig(); + config.FlowConfig = TDeviceOidcConfig{"public-client", {}}; + config.Acceptor(std::make_shared()); + auto factory = CreateOidcProviderFactory(config); + auto provider = factory->CreateProvider(); + auto result = provider->GetAuthInfoAsync(); + const bool completed = result.Wait(TDuration::Seconds(4)); + gate.TrySetValue(); + UNIT_ASSERT(completed); + UNIT_ASSERT(result.HasException()); +} + +Y_UNIT_TEST(DiscardedDeliveryCompletesOnFacilityDestruction) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + TStaticOidcConfig flow; + flow.AccessToken = "static"; + config.FlowConfig = flow; + config.Cacher(cache); + auto factory = CreateOidcProviderFactory(config); + auto provider = factory->CreateProvider(facility); + auto pending = provider->GetAuthInfoAsync(); + cache->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + facility.reset(); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(1))); + UNIT_ASSERT(pending.HasException()); +} + +Y_UNIT_TEST(QueuedDeliveryNeverReturnsExpiredToken) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + TStaticOidcConfig flow; + flow.AccessToken = "static"; + flow.ExpiresAt = TInstant::Now() + TDuration::Seconds(1); + config.FlowConfig = flow; + config.Cacher(cache); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + auto pending = provider->GetAuthInfoAsync(); + cache->Release.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + NThreading::NewPromise().GetFuture().Wait(*flow.ExpiresAt + TDuration::MilliSeconds(1)); + facility->RunTasks(); + UNIT_ASSERT(pending.HasException()); +} +Y_UNIT_TEST(ResponseQueueFailureCompletesPendingFuture) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{"access", std::nullopt}; + config.Cacher(cache); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + auto pending = provider->GetAuthInfoAsync(); + cache->Release.TrySetValue(); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(5))); + UNIT_ASSERT_EXCEPTION_CONTAINS(pending.GetValueSync(), std::exception, "response queue unavailable"); +} + +Y_UNIT_TEST(RejectsMissingRequiredCredentialsAndInvalidScopes) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + for (const TFlowConfig& flow : { + TFlowConfig{TStaticOidcConfig{}}, + TFlowConfig{TClientOidcConfig{"", "secret", {}}}, + TFlowConfig{TDeviceOidcConfig{"", {}}}}) { + config.FlowConfig = flow; + UNIT_ASSERT_EXCEPTION(CreateOidcProviderFactory(config), std::invalid_argument); + } + for (const std::string& scope : {"", "read write", "read\twrite", "read\nwrite", "read\"", "read\\", "read\x7f", "профиль"}) { + config.FlowConfig = TClientOidcConfig{"client", "secret", {scope}}; + UNIT_ASSERT_EXCEPTION(CreateOidcProviderFactory(config), std::invalid_argument); + } +} + +Y_UNIT_TEST(ClientIdentityNormalizesScopesButDistinguishesCredentials) { + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TClientOidcConfig{"client", "secret", {"write", "read", "read"}}; + const auto identity = GetOidcClientIdentity(config); + auto& client = std::get(config.FlowConfig); + client.Scopes = {"read", "write"}; + UNIT_ASSERT_VALUES_EQUAL(GetOidcClientIdentity(config), identity); + client.ClientSecret = "other-secret"; + UNIT_ASSERT(GetOidcClientIdentity(config) != identity); + config.FlowConfig = TStaticOidcConfig{"token", std::nullopt}; + const auto staticIdentity = GetOidcClientIdentity(config); + std::get(config.FlowConfig).ExpiresAt = TInstant::Seconds(100); + UNIT_ASSERT(GetOidcClientIdentity(config) != staticIdentity); + const auto factory = CreateOidcProviderFactory(config); + UNIT_ASSERT(factory->CreateProvider() == factory->CreateProvider()); +} + +Y_UNIT_TEST(CachedTokenWithoutExpiryNeedsNoRefresh) { + TOidcTestServer server; + auto cache = std::make_shared(); + cache->Write({{"cached", std::nullopt}, std::nullopt}); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer cached"); + UNIT_ASSERT(provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); +} + +Y_UNIT_TEST(UnknownCachedLifetimeTriggersRefresh) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"unknown-expiry", std::nullopt}, TOAuthToken{"refresh", std::nullopt}}); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer fresh"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); + UNIT_ASSERT_VALUES_EQUAL(server.Requests()[0].Form.Get("grant_type"), "refresh_token"); +} + +Y_UNIT_TEST(ExpiredRefreshTokenTriggersNewClientGrant) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"fresh","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"expired-refresh", TInstant::Seconds(1)}}); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer fresh"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); + UNIT_ASSERT_VALUES_EQUAL(server.Requests()[0].Form.Get("grant_type"), "client_credentials"); +} + +Y_UNIT_TEST(TerminalRefreshErrorDoesNotStartNewGrant) { + TOidcTestServer server; + server.Enqueue(R"({"error":"invalid_client"})", HTTP_UNAUTHORIZED); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"refresh", std::nullopt}}); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + auto result = provider->GetAuthInfoAsync(); + UNIT_ASSERT(result.Wait(TDuration::Seconds(5))); + UNIT_ASSERT_EXCEPTION_CONTAINS(result.GetValueSync(), std::exception, "invalid_client"); + UNIT_ASSERT(!provider->IsValid()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); +} + +Y_UNIT_TEST(StopDuringCacheWriteDoesNotPublishToken) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{"access", std::nullopt}; + config.Cacher(cache); + auto provider = std::make_shared(config, facility); + auto pending = provider->GetAuthInfoAsync(); + const bool entered = cache->Entered.GetFuture().Wait(TDuration::Seconds(5)); + auto stopped = std::async(std::launch::async, [provider] { provider->Stop(); }); + const bool cancelled = pending.Wait(TDuration::Seconds(5)); + cache->Release.TrySetValue(); + stopped.get(); + UNIT_ASSERT(entered); + UNIT_ASSERT(cancelled); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT(provider->GetAuthInfoAsync().HasException()); + UNIT_ASSERT(!provider->IsValid()); +} + +Y_UNIT_TEST(FacilityExpiredDuringCacheWriteCompletesWithError) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{"access", std::nullopt}; + config.Cacher(cache); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + const auto pending = provider->GetAuthInfoAsync(); + const bool entered = cache->Entered.GetFuture().Wait(TDuration::Seconds(5)); + facility.reset(); + cache->Release.TrySetValue(); + UNIT_ASSERT(entered); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(5))); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT(!provider->IsValid()); +} + +Y_UNIT_TEST(StopIgnoresThrowingPendingSubscriber) { + auto cache = std::make_shared(); + auto facility = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{"access", std::nullopt}; + config.Cacher(cache); + NOidc::NPrivate::TStaticProvider provider(config, facility); + auto pending = provider.GetAuthInfoAsync(); + pending.Subscribe([](const auto&) { throw std::runtime_error("subscriber failure"); }); + auto stopped = std::async(std::launch::async, [&] { provider.Stop(); }); + const bool cancelled = pending.Wait(TDuration::Seconds(5)); + cache->Release.TrySetValue(); + UNIT_ASSERT_NO_EXCEPTION(stopped.get()); + UNIT_ASSERT(cancelled); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT(!provider.IsValid()); +} + +Y_UNIT_TEST(StopIgnoresThrowingQueuedSubscriber) { + auto cache = std::make_shared(); + TOidcConfig config; + config.Issuer = "https://issuer.example"; + config.FlowConfig = TStaticOidcConfig{"access", std::nullopt}; + config.Cacher(cache); + auto facility = std::make_shared(); + NOidc::NPrivate::TStaticProvider provider(config, facility); + auto pending = provider.GetAuthInfoAsync(); + pending.Subscribe([](const auto&) { throw std::runtime_error("subscriber failure"); }); + cache->Release.TrySetValue(); + const bool queued = facility->WaitForTask(); + std::exception_ptr stopError; + try { + provider.Stop(); + } catch (...) { + stopError = std::current_exception(); + } + UNIT_ASSERT(queued); + UNIT_ASSERT(stopError == nullptr); + UNIT_ASSERT(pending.HasException()); + UNIT_ASSERT_NO_EXCEPTION(facility->RunTasks()); +} + +Y_UNIT_TEST(TransientInvalidGrantRetriesRefreshWithoutStartingAnotherFlow) { + for (const bool device : {false, true}) { + TOidcTestServer server; + auto replyGate = NThreading::NewPromise(); + server.BlockTokenRepliesUntil(replyGate.GetFuture()); + server.Enqueue(R"({"error":"invalid_grant"})", HTTP_SERVICE_UNAVAILABLE); + server.Enqueue(R"({"access_token":"refreshed","token_type":"Bearer","expires_in":600})", HTTP_OK); + auto cache = std::make_shared(); + cache->Write({{"expired", TInstant::Seconds(1)}, TOAuthToken{"refresh", std::nullopt}}); + auto config = server.ClientConfig().Cacher(cache); + if (device) { + config.FlowConfig = TDeviceOidcConfig{"client", {}}; + } + auto facility = std::make_shared(); + auto provider = CreateOidcProviderFactory(config)->CreateProvider(facility); + auto firstAttempt = provider->GetAuthInfoAsync(); + replyGate.TrySetValue(); + UNIT_ASSERT(facility->WaitForTask()); + facility->RunTasks(); + UNIT_ASSERT(firstAttempt.HasException()); + if (!provider->GetAuthInfoAsync().HasValue()) { + UNIT_ASSERT(facility->WaitForTask()); + facility->RunTasks(); + } + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer refreshed"); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 2); + for (const auto& request : requests) { + UNIT_ASSERT_VALUES_EQUAL(request.Form.Get("grant_type"), "refresh_token"); + UNIT_ASSERT_VALUES_EQUAL(request.Form.Get("refresh_token"), "refresh"); + } + } +} + +Y_UNIT_TEST(PersistentOutageSettlesPendingCredentials) { + TOidcTestServer server; + for (size_t i = 0; i < 10; ++i) { + server.Enqueue(R"({"error":"temporarily_unavailable"})", HTTP_SERVICE_UNAVAILABLE); + } + auto provider = CreateOidcProviderFactory(server.ClientConfig())->CreateProvider(); + auto pending = provider->GetAuthInfoAsync(); + UNIT_ASSERT(pending.Wait(TDuration::Seconds(2))); + UNIT_ASSERT_EXCEPTION_CONTAINS(pending.GetValueSync(), std::exception, "503"); + UNIT_ASSERT_EXCEPTION_CONTAINS(provider->GetAuthInfo(), std::exception, "503"); + UNIT_ASSERT(!provider->IsValid()); +} + +Y_UNIT_TEST(ShutdownJoinsWorkerDuringTlsHandshake) { + TOidcTestServer server; + auto gate = NThreading::NewPromise(); + server.BlockTlsHandshakeUntil(gate.GetFuture()); + auto provider = CreateOidcProviderFactory(server.ClientConfig())->CreateProvider(); + auto pending = provider->GetAuthInfoAsync(); + const bool started = server.WaitForTlsHandshake(); + auto stopped = std::async(std::launch::async, [provider = std::move(provider)]() mutable { + provider.reset(); + }); + const bool completed = stopped.wait_for(std::chrono::seconds(2)) == std::future_status::ready; + gate.TrySetValue(); + stopped.get(); + UNIT_ASSERT(started); + UNIT_ASSERT(!completed); + UNIT_ASSERT(pending.HasException()); +} + +Y_UNIT_TEST(ClientAcceptsOpaqueTokenWithoutLifetime) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"opaque","token_type":"Bearer"})", HTTP_OK); + auto cache = std::make_shared(); + auto provider = CreateOidcProviderFactory(server.ClientConfig().Cacher(cache))->CreateProvider(); + UNIT_ASSERT_VALUES_EQUAL(provider->GetAuthInfo(), "Bearer opaque"); + UNIT_ASSERT(provider->IsValid()); + const auto stored = cache->Read(); + UNIT_ASSERT(stored.has_value() && !stored->AccessToken.ExpiresAt.has_value()); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); +} + +} // Y_UNIT_TEST_SUITE(TOidcCredentials) diff --git a/tests/unit/client/oidc/helpers/test_server.cpp b/tests/unit/client/oidc/helpers/test_server.cpp new file mode 100644 index 00000000000..850bef49b2b --- /dev/null +++ b/tests/unit/client/oidc/helpers/test_server.cpp @@ -0,0 +1,329 @@ +#include "test_server.h" + +#include +#include +#include + +#include +#include + +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +namespace { + +// Self-signed localhost certificate and private key used only by the test server. +constexpr char TestCertificate[] = R"PEM(-----BEGIN CERTIFICATE----- +MIIDITCCAgmgAwIBAgIUUwx3TvJXvw6bLn1TDL4NXW6haQ8wDQYJKoZIhvcNAQEL +BQAwFDESMBAGA1UEAwwJbG9jYWxob3N0MCAXDTI2MDkxNTEwMjUxM1oYDzIxMjYw +ODIyMTAyNTEzWjAUMRIwEAYDVQQDDAlsb2NhbGhvc3QwggEiMA0GCSqGSIb3DQEB +AQUAA4IBDwAwggEKAoIBAQCoEDP1zRgNbHt5nh5CmB/4f5Ic+znXohPE3nSpfxro +njCTcLmOYYN34yxij7yozebAMJV5ebdJIhy4H90nCQpXV448aEf3cVpl9s0kEdDG +t0GFp31AtK+U3SzAERmx+NG8hyA2GHYRYgirglpueVtDtiRu8Um2Uu2M+ieHXgJx +sn4bOgo1Y+9r1SMSJXJtbC2SMcxyJC4HjzORsse0+cazZ66Ey9Jw9fQ8gKJs0JjA +q5MKno6a/BMrXu1JDlY3FOFKM6saIRKfH6GtmrM6EYmwSdC2CjBXLKI/BDrzDPQb +CsMSYjlxkiykmmDFGQuSYEMOm5QO33NEGGSWBAtUfSxdAgMBAAGjaTBnMB0GA1Ud +DgQWBBTGCtXyTRdpK5HQDS0SyeCTQ9xeKzAfBgNVHSMEGDAWgBTGCtXyTRdpK5HQ +DS0SyeCTQ9xeKzAUBgNVHREEDTALgglsb2NhbGhvc3QwDwYDVR0TAQH/BAUwAwEB +/zANBgkqhkiG9w0BAQsFAAOCAQEAV2ajayHUpaXRW876j8Vfa4AueSa3buYaXrzc +d8aKrlpEcstVlCykhBIHnPzlWXqfTbkNBYer9C/xfyXJE9m6xTW1OIQPibq3iuRi +6+vz49gsLyhudoP3gz3oHM1of+5YEp3vh4gzxTojS19ffLlgBWEUTuuETNdAakW4 +++4Q7eus0GUlrrZTZeqjnhzU1UudjKc1ntXBZTzOAsN3Vt0BrO/eFTix7MDI3ktK +wGPQRUTNm0E1ITQ+Vst4HDouCZI34zdOSsJvZouWzo6lSiRNm+KUn6cqay9cNxoQ +36XwMgo6i+KWJSo2dVJXGaQk51u1jptkCXKvYJo2xq8sLdXBfA== +-----END CERTIFICATE----- +)PEM"; + +constexpr char TestPrivateKey[] = R"PEM(-----BEGIN PRIVATE KEY----- +MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQCoEDP1zRgNbHt5 +nh5CmB/4f5Ic+znXohPE3nSpfxronjCTcLmOYYN34yxij7yozebAMJV5ebdJIhy4 +H90nCQpXV448aEf3cVpl9s0kEdDGt0GFp31AtK+U3SzAERmx+NG8hyA2GHYRYgir +glpueVtDtiRu8Um2Uu2M+ieHXgJxsn4bOgo1Y+9r1SMSJXJtbC2SMcxyJC4HjzOR +sse0+cazZ66Ey9Jw9fQ8gKJs0JjAq5MKno6a/BMrXu1JDlY3FOFKM6saIRKfH6Gt +mrM6EYmwSdC2CjBXLKI/BDrzDPQbCsMSYjlxkiykmmDFGQuSYEMOm5QO33NEGGSW +BAtUfSxdAgMBAAECggEAALFcaCQpzWMH7pw/wgTa2405vs6Bp7QT18kWUF06m6s3 +RmGoaiqtvjtHWLqrS5jZstYgb56YVQAD1Uslqr5coWKLhA/mqAxQE+vctEwx1k01 +bcXJ0Y/3yf8launxzKwwFSe2HXL5XaClVNZV5W8GTh+nQ8vRLXlnCvXnCXr85ezA +O4APlFi1qbpbNY9wzVXYOIfocq+CTEvQASOcvvwLAk1Tuh3Vcibl7xRFWlXkguea +LBSfvIRgVzPomMTgyJsve1W+w09C86I8yawa3zWiFS8O5HdCy8T7CvTCEQKA/hPI +fv7nDCsIuqjXZsuLw6PJN0sBiN18C/ARrlcH0wBPYQKBgQDbeGx6PAdfOBe4nUte +MJPQWBLnhLyWB44NOQA3FLoFW9VEu51Aa+7I/nOUxFFj29FxSJxkftvQu92laLow +o+sdXnjGfWd4yh/MeC7DOXOHmtX+PDs/FypPZMUbwawNkYwog16nd/z8BuBIsQkO +2bvfOBpDqA/wOVNCFawaCJH4rQKBgQDECVXvm5hAMTHeI2CbAQJ6PuT0eTYFJ4oa +jDqmK5UzMkxshinF/cDPWabwVz9kJffoQc6JrhoTxaewaeg7iqoTp5s2/5BzPTXC +qKb13RQtNqQ7yl0lBXjJ/VZQ2BVAUh3h/kpwRP6f7eMY5YYPfxvrbCru7Lb1dHi+ +TcOmuY0IcQKBgQCYJItO0X5qy//lw2UUDqjpraStSp9Rgjs/f1xe0seCH39g/o6s +siX+wCZv4whpKWGwHp4MLMVFlna4zDkGrxu2aF9hel3YpoYUwNvqClHEl9nxPN/1 +hKGYGEtsSn5ziYqYKzna7ps6O6oPumqFGPvcapAKht9FsPe+wDdmdLp8oQKBgH5h +Ll+cRZkMngOBdyQ2kGxS47Of+O11whi/UogSDMvGn3JPQ9r6bjS+rVrARIPB3oKC ++i3UacdZY3PdsvO/v0mQggYA2BUS3vexVoGmlv1W/qX1Hfth/a7qfZz80SZ4Sf+J +ul+Ke0SLTh6cycJvxYYOY9dID+NJxRWaeImhkYRhAoGABGWTKewptIvJT8n69Rmt +hiaSpksnRsjT/URNRD9oyGPeFbmh3BjoUJND/FhMwMr/y446Yo9n6+2Wi9U52n3I +Bmc27AR902ceiZY5oPAoe82wGTBqTMvfeLZSNGZ67WxDqNH/2Vy52/21IM/KAEVk +2GMqMlo0l4Zn5vQ9jTkfKBI= +-----END PRIVATE KEY----- +)PEM"; + +SSL_CTX* ServerContext(); + +class TTlsStreams: public THttpServerConn::ISocketStreams, public IInputStream, public IOutputStream { +public: + explicit TTlsStreams(const TSocket& socket); + IInputStream* Input() override; + IOutputStream* Output() override; + void Reset() override; + +private: + size_t DoRead(void* buffer, size_t size) override; + void DoWrite(const void* buffer, size_t size) override; + + TSocket Socket; + std::unique_ptr Ssl; +}; + +SSL_CTX* ServerContext() { + static const auto context = [] { + std::unique_ptr certificateBio( + BIO_new_mem_buf(TestCertificate, sizeof(TestCertificate) - 1), BIO_free); + std::unique_ptr keyBio( + BIO_new_mem_buf(TestPrivateKey, sizeof(TestPrivateKey) - 1), BIO_free); + Y_ENSURE(certificateBio != nullptr && keyBio != nullptr); + std::unique_ptr certificate( + PEM_read_bio_X509(certificateBio.get(), nullptr, nullptr, nullptr), X509_free); + std::unique_ptr key( + PEM_read_bio_PrivateKey(keyBio.get(), nullptr, nullptr, nullptr), EVP_PKEY_free); + Y_ENSURE(certificate != nullptr && key != nullptr); + std::unique_ptr result(SSL_CTX_new(TLS_server_method()), SSL_CTX_free); + Y_ENSURE(result != nullptr, "Cannot create test TLS context"); + Y_ENSURE(SSL_CTX_use_certificate(result.get(), certificate.get()) == 1); + Y_ENSURE(SSL_CTX_use_PrivateKey(result.get(), key.get()) == 1); + Y_ENSURE(SSL_CTX_check_private_key(result.get()) == 1); + + // The HTTP client loads trusted CAs from SSL_CERT_FILE. Keep this + // temporary copy alive for the test process; TTempFileHandle removes it. + static TTempFileHandle trustedCertificate; + trustedCertificate.Write(TestCertificate, sizeof(TestCertificate) - 1); + trustedCertificate.Close(); + SetEnv("SSL_CERT_FILE", trustedCertificate.Name()); + return result; + }(); + return context.get(); +} + +TTlsStreams::TTlsStreams(const TSocket& socket) + : Socket(socket) + , Ssl(SSL_new(ServerContext()), SSL_free) +{ + Y_ENSURE(Ssl != nullptr); + Y_ENSURE(SSL_set_fd(Ssl.get(), Socket) == 1); + Y_ENSURE(SSL_accept(Ssl.get()) == 1, "Test TLS handshake failed"); +} + +IInputStream* TTlsStreams::Input() { + return this; +} + +IOutputStream* TTlsStreams::Output() { + return this; +} + +void TTlsStreams::Reset() { +} + +size_t TTlsStreams::DoRead(void* buffer, size_t size) { + const int result = SSL_read(Ssl.get(), buffer, std::min(size, std::numeric_limits::max())); + if (result <= 0 && SSL_get_error(Ssl.get(), result) == SSL_ERROR_ZERO_RETURN) { + return 0; + } + Y_ENSURE(result > 0, "Test TLS read failed"); + return result; +} + +void TTlsStreams::DoWrite(const void* buffer, size_t size) { + auto data = static_cast(buffer); + while (size) { + const int written = SSL_write(Ssl.get(), data, std::min(size, std::numeric_limits::max())); + Y_ENSURE(written > 0, "Test TLS write failed"); + data += written; + size -= written; + } +} + +} // namespace + +TOidcTestServer::TOidcTestServer() + : Options(Ports.GetPort()) + , Server(this, Options) +{ + ServerContext(); + Y_ENSURE(Server.Start(), "Cannot start test OIDC server"); +} + +TOidcTestServer::~TOidcTestServer() { + Server.Stop(); +} + +std::vector> TOidcTestServer::HostHeaders() const { + with_lock (Mutex) { + return RecordedHosts; + } +} + +std::string TOidcTestServer::Issuer() const { + return "https://localhost:" + std::to_string(Options.Port) + "/realm"; +} + +NYdb::NOidc::TOidcConfig TOidcTestServer::ClientConfig() const { + NYdb::NOidc::TOidcConfig config; + config.Issuer = Issuer(); + config.FlowConfig = NYdb::NOidc::TClientOidcConfig{"client", "secret +&", {"read", "write"}}; + return config; +} + +void TOidcTestServer::Enqueue(TString body, HttpCodes status) { + with_lock (Mutex) { + Replies.push_back({status, std::move(body)}); + } +} + +void TOidcTestServer::SetDiscoveryReply(TString body, HttpCodes status) { + with_lock (Mutex) { + DiscoveryReply = TReply{status, std::move(body)}; + } +} + +void TOidcTestServer::SetTokenReplyDelay(TDuration delay) { + with_lock (Mutex) { + TokenReplyDelay = delay; + } +} + +void TOidcTestServer::BlockTokenRepliesUntil(NThreading::TFuture released) { + with_lock (Mutex) { + TokenReplyGate = std::move(released); + } +} + +std::vector TOidcTestServer::Requests() const { + with_lock (Mutex) { + return Recorded; + } +} + +void TOidcTestServer::BlockTlsHandshakeUntil(NThreading::TFuture released) { + with_lock (Mutex) { + TlsHandshakeGate = std::move(released); + } +} + +bool TOidcTestServer::WaitForTlsHandshake() { + with_lock (Mutex) { + return Changed.wait_for(Mutex, std::chrono::seconds(10), [&] { return TlsHandshakeStarted; }); + } +} + +size_t TOidcTestServer::DiscoveryCount() const { + with_lock (Mutex) { + return Discoveries; + } +} + +bool TOidcTestServer::WaitRequests(size_t count) { + with_lock (Mutex) { + return Changed.wait_for(Mutex, std::chrono::seconds(10), [&] { return Recorded.size() >= count; }); + } +} + +TOidcTestServer::TRequest::TRequest(TOidcTestServer& server) + : Server(server) +{ +} + +bool TOidcTestServer::TRequest::DoReply(const TReplyParams& params) { + const TParsedHttpFull parsed(params.Input.FirstLine()); + const TString body = params.Input.ReadAll(); + TReply reply; + NThreading::TFuture replyGate; + TDuration replyDelay; + with_lock (Server.Mutex) { + std::vector hosts; + for (const auto& header : params.Input.Headers()) { + if (header.Name() == "Host") { + hosts.push_back(header.Value()); + } + } + Server.RecordedHosts.push_back(std::move(hosts)); + if (parsed.Path == "/realm/.well-known/openid-configuration") { + ++Server.Discoveries; + NJson::TJsonValue metadata; + metadata["issuer"] = Server.Issuer(); + metadata["token_endpoint"] = Server.Issuer() + "/token"; + metadata["device_authorization_endpoint"] = Server.Issuer() + "/device"; + reply.Body = NJson::WriteJson(metadata, false); + if (Server.DiscoveryReply.has_value()) { + reply = *Server.DiscoveryReply; + } + } else { + TRequestInfo request{TString(parsed.Path), TString(parsed.Method), TCgiParameters(body), {}}; + for (const auto& header : params.Input.Headers()) { + if (header.Name() == "Authorization") { + request.Authorization = header.Value(); + } + } + Server.Recorded.push_back(std::move(request)); + if (parsed.Path == "/realm/token") { + replyGate = Server.TokenReplyGate; + replyDelay = Server.TokenReplyDelay; + } + if (Server.Replies.empty()) { + reply = {HTTP_BAD_REQUEST, R"({"error":"unexpected_request"})"}; + } else { + reply = std::move(Server.Replies.front()); + Server.Replies.pop_front(); + } + Server.Changed.notify_all(); + } + } + if (replyGate.Initialized()) { + replyGate.Wait(); + } + if (replyDelay) { + params.Output << "HTTP/1.1 " << static_cast(reply.Status) + << " OK\r\nContent-Length: " << reply.Body.size() << "\r\n\r\n"; + for (const char c : reply.Body) { + params.Output.Write(c); + params.Output.Flush(); + Sleep(replyDelay); + } + return true; + } + THttpResponse response(reply.Status); + response.SetContent(reply.Body); + response.OutTo(params.Output); + return true; +} + +TClientRequest* TOidcTestServer::CreateClient() { + return new TRequest(*this); +} + +THolder TOidcTestServer::TRequest::CreateHttpConnection(const TSocket& socket, size_t outputBuffer) { + NThreading::TFuture gate; + with_lock (Server.Mutex) { + gate = Server.TlsHandshakeGate; + Server.TlsHandshakeStarted = true; + } + Server.Changed.notify_all(); + if (gate.Initialized()) { + gate.Wait(); + } + return MakeHolder(MakeHolder(socket), outputBuffer); +} diff --git a/tests/unit/client/oidc/helpers/test_server.h b/tests/unit/client/oidc/helpers/test_server.h new file mode 100644 index 00000000000..998689c5202 --- /dev/null +++ b/tests/unit/client/oidc/helpers/test_server.h @@ -0,0 +1,87 @@ +#pragma once + +#include + +#include +#include +#include +#include + +#include + +#include +#include +#include + +class TOidcTestServer: public THttpServer::ICallBack { +public: + struct TReply { + HttpCodes Status = HTTP_OK; + TString Body; + }; + + struct TRequestInfo { + TString Path; + TString Method; + TCgiParameters Form; + TString Authorization; + }; + + TOidcTestServer(); + + ~TOidcTestServer() override; + + std::string Issuer() const; + + NYdb::NOidc::TOidcConfig ClientConfig() const; + + void Enqueue(TString body, HttpCodes status); + + void SetDiscoveryReply(TString body, HttpCodes status); + + void SetTokenReplyDelay(TDuration delay); + + void BlockTokenRepliesUntil(NThreading::TFuture released); + + void BlockTlsHandshakeUntil(NThreading::TFuture released); + + bool WaitForTlsHandshake(); + + std::vector Requests() const; + + std::vector> HostHeaders() const; + + size_t DiscoveryCount() const; + + bool WaitRequests(size_t count); + + class TRequest: public TRequestReplier { + public: + explicit TRequest(TOidcTestServer& server); + + bool DoReply(const TReplyParams& params) override; + + THolder CreateHttpConnection(const TSocket& socket, size_t outputBuffer) override; + + private: + TOidcTestServer& Server; + }; + + TClientRequest* CreateClient() override; + +private: + TPortManager Ports; + THttpServer::TOptions Options; + mutable TMutex Mutex; + std::condition_variable_any Changed; + std::deque Replies; + std::optional DiscoveryReply; + std::vector Recorded; + std::vector> RecordedHosts; + size_t Discoveries = 0; + NThreading::TFuture TokenReplyGate; + NThreading::TFuture TlsHandshakeGate; + bool TlsHandshakeStarted = false; + TDuration TokenReplyDelay; + THttpServer Server; +}; diff --git a/tests/unit/client/oidc/protocol_ut.cpp b/tests/unit/client/oidc/protocol_ut.cpp new file mode 100644 index 00000000000..37c4b1f8a69 --- /dev/null +++ b/tests/unit/client/oidc/protocol_ut.cpp @@ -0,0 +1,422 @@ +#include "test_server.h" + +#include +#include + +#include +#include +#include +#include + +#include +#include + +using namespace NYdb; +using namespace NYdb::NOidc; +using namespace NYdb::NOidc::NPrivate; + +namespace { + +NJson::TJsonValue Json(const TString& text); +NJson::TJsonValue Metadata(const TOidcTestServer& server); +NJson::TJsonValue DeviceResponse(const TOidcTestServer& server); +TOidcConfig DeviceConfig(const TOidcTestServer& server); +std::string Jwt(const TString& payload); + +NJson::TJsonValue Json(const TString& text) { + NJson::TJsonValue result; + UNIT_ASSERT(NJson::ReadJsonTree(text, &result)); + return result; +} + +NJson::TJsonValue Metadata(const TOidcTestServer& server) { + NJson::TJsonValue result; + result["issuer"] = server.Issuer(); + result["token_endpoint"] = server.Issuer() + "/token"; + result["device_authorization_endpoint"] = server.Issuer() + "/device"; + return result; +} + +NJson::TJsonValue DeviceResponse(const TOidcTestServer& server) { + auto result = Json(R"({"device_code":"private-code","user_code":"ABCD","expires_in":600})"); + result["verification_uri"] = server.Issuer() + "/verify"; + return result; +} + +TOidcConfig DeviceConfig(const TOidcTestServer& server) { + auto config = server.ClientConfig().Acceptor(std::make_shared()); + config.FlowConfig = TDeviceOidcConfig{"public-client", {"read"}}; + return config; +} + +std::string Jwt(const TString& payload) { + return "e30." + std::string(Base64EncodeUrl(payload)) + ".signature"; +} + +} // namespace + +Y_UNIT_TEST_SUITE(TOidcProtocol) { +Y_UNIT_TEST(DiscoveryAndTokenRequestsIncludePortInHost) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_VALUES_EQUAL(protocol.ClientGrant().AccessToken.Token, "access"); + // Issuer is https://localhost:/realm. + const auto expectedHost = config.Issuer.substr(8, config.Issuer.size() - 8 - 6); + const auto requests = server.HostHeaders(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 2); + for (const auto& hosts : requests) { + UNIT_ASSERT_VALUES_EQUAL(hosts.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(hosts.front(), expectedHost); + } +} + +Y_UNIT_TEST(RejectsInvalidDurations) { + for (const TString& value : {"-1", "1.5", "true", "null", "\"10\"", "18446744073709551615"}) { + UNIT_ASSERT_EXCEPTION(Seconds(Json(value), "expires_in", false), TError); + } + UNIT_ASSERT_EXCEPTION(Seconds(Json("0"), "expires_in", false), TError); + UNIT_ASSERT_VALUES_EQUAL(Seconds(Json("0"), "refresh_expires_in", true), 0); + const ui64 maximum = std::numeric_limits::max() / 2 / 1'000'000; + UNIT_ASSERT_VALUES_EQUAL(Seconds(NJson::TJsonValue(maximum), "expires_in", false), maximum); + UNIT_ASSERT_EXCEPTION(Seconds(NJson::TJsonValue(maximum + 1), "expires_in", false), TError); +} + +Y_UNIT_TEST(JwtExpiryHandlesMissingAndInvalidPayloads) { + for (const auto& token : {std::string("header.payload"), std::string("header.!.signature"), + Jwt("[]"), Jwt("{}"), Jwt("null"), Jwt("not-json"), + "header." + std::string(1024 * 1024, 'a') + ".signature"}) { + UNIT_ASSERT(!JwtExpiry(token).has_value()); + } + UNIT_ASSERT_VALUES_EQUAL(*JwtExpiry(Jwt(R"({"exp":0})")), TInstant::Zero()); + UNIT_ASSERT_VALUES_EQUAL(*JwtExpiry(Jwt(R"({"exp":-1})")), TInstant::Zero()); + for (const auto& payload : {R"({"exp":"tomorrow"})", R"({"exp":null})", + R"({"exp":true})", R"({"exp":18446744073709551615})"}) { + UNIT_ASSERT(!JwtExpiry(Jwt(payload)).has_value()); + } +} + +Y_UNIT_TEST(UrlErrorsIdentifyIssuerOrEndpoint) { + for (const bool issuer : {false, true}) { + const std::string role = issuer ? "issuer" : "endpoint"; + UNIT_ASSERT_EXCEPTION_CONTAINS(ParseUrl("https://user:secret@example.com", issuer), + std::invalid_argument, "invalid " + role + " URL"); + UNIT_ASSERT_EXCEPTION_CONTAINS(ParseUrl("http://example.com", issuer), + std::invalid_argument, role + " requires HTTPS"); + } +} + +Y_UNIT_TEST(DiscoveryMismatchReportsIssuerIdentifiers) { + TOidcTestServer server; + auto metadata = Metadata(server); + metadata["issuer"] = server.Issuer() + "/"; + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.ClientGrant(), TError, + "configured '" + config.Issuer + "', advertised '" + config.Issuer + "/'"); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(InvalidAdvertisedIssuerDoesNotLeakSecrets) { + for (const auto& issuer : {"https://user:private-token@example.com", + "https://example.com?token=private-token", "https://example.com/#private-token"}) { + TOidcTestServer server; + auto metadata = Metadata(server); + metadata["issuer"] = issuer; + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + try { + protocol.ClientGrant(); + UNIT_FAIL("expected an invalid advertised issuer error"); + } catch (const TError& error) { + const std::string message = error.what(); + UNIT_ASSERT_STRING_CONTAINS(message, "discovery issuer mismatch: invalid advertised issuer URL"); + UNIT_ASSERT(message.find("private-token") == std::string::npos); + } + UNIT_ASSERT(server.Requests().empty()); + } +} + +Y_UNIT_TEST(RejectsInvalidDiscoveryFields) { + for (const auto& field : {"issuer", "token_endpoint"}) { + for (const TString& value : {"null", "false", "\"\""}) { + TOidcTestServer server; + auto metadata = Metadata(server); + metadata[field] = Json(value); + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION(protocol.ClientGrant(), TError); + UNIT_ASSERT(server.Requests().empty()); + } + } + TOidcTestServer server; + auto metadata = Metadata(server); + metadata["issuer"] = "https://another-issuer.example"; + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.ClientGrant(), TError, "issuer mismatch"); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(DiscoveryRejectsUnsupportedAuthenticationMethods) { + for (const TString& methods : {"null", "\"client_secret_basic\"", "[42]", "[]", "[\"client_secret_post\"]"}) { + TOidcTestServer server; + auto metadata = Metadata(server); + metadata["token_endpoint_auth_methods_supported"] = Json(methods); + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION(protocol.ClientGrant(), TError); + UNIT_ASSERT(server.Requests().empty()); + } +} + +Y_UNIT_TEST(DiscoveryAcceptsBasicAuthAndTrailingIssuerSlash) { + TOidcTestServer server; + auto config = server.ClientConfig(); + config.Issuer += "/"; + auto metadata = Metadata(server); + metadata["issuer"] = config.Issuer; + metadata["token_endpoint_auth_methods_supported"] = Json("[\"client_secret_post\",\"client_secret_basic\"]"); + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + server.Enqueue(R"({"access_token":"access","token_type":"Bearer","expires_in":600})", HTTP_OK); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_VALUES_EQUAL(protocol.ClientGrant().AccessToken.Token, "access"); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 1); +} + +Y_UNIT_TEST(RejectsMissingAndInvalidTokenFields) { + for (const auto& field : {"access_token", "token_type", "refresh_token"}) { + for (const TString& value : {"null", "42", "\"\"", "\"bad token\"", "\"bad\\u007ftoken\""}) { + TOidcTestServer server; + auto response = Json(R"({"access_token":"access","token_type":"Bearer","expires_in":600})"); + response[field] = Json(value); + server.Enqueue(NJson::WriteJson(response, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION(protocol.ClientGrant(), TError); + } + } + for (const TString& response : {"{}", "[]", R"({"token_type":"Bearer","expires_in":600})"}) { + TOidcTestServer server; + server.Enqueue(response, HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION(protocol.ClientGrant(), TError); + } +} + +Y_UNIT_TEST(AccessTokenCanUseJwtExpiry) { + TOidcTestServer server; + const auto expiry = TInstant::Now() + TDuration::Hours(1); + NJson::TJsonValue payload; + payload["exp"] = expiry.Seconds(); + const auto token = Jwt(NJson::WriteJson(payload, false)); + auto response = Json(R"({"token_type":"Bearer"})"); + response["access_token"] = token; + server.Enqueue(NJson::WriteJson(response, false), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + const auto result = protocol.ClientGrant(); + UNIT_ASSERT_VALUES_EQUAL(result.AccessToken.Token, token); + UNIT_ASSERT(result.AccessToken.ExpiresAt.has_value()); + UNIT_ASSERT_VALUES_EQUAL(result.AccessToken.ExpiresAt->Seconds(), expiry.Seconds()); +} + +Y_UNIT_TEST(RefreshCanRemoveExpiryAndRotateToken) { + TOidcTestServer server; + server.Enqueue(R"({"access_token":"first","token_type":"Bearer","expires_in":600,"refresh_expires_in":0})", HTTP_OK); + server.Enqueue(R"({"access_token":"second","token_type":"Bearer","expires_in":600,"refresh_token":"rotated","refresh_expires_in":1200})", HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + const auto first = protocol.Refresh({"refresh", TInstant::Now() + TDuration::Minutes(1)}); + UNIT_ASSERT(first.RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(first.RefreshToken->Token, "refresh"); + UNIT_ASSERT(!first.RefreshToken->ExpiresAt.has_value()); + const auto second = protocol.Refresh(*first.RefreshToken); + UNIT_ASSERT(second.RefreshToken.has_value() && second.RefreshToken->ExpiresAt.has_value()); + UNIT_ASSERT_VALUES_EQUAL(second.RefreshToken->Token, "rotated"); + UNIT_ASSERT(second.RefreshToken->IsValid(TInstant::Now())); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 1); +} + +Y_UNIT_TEST(HttpErrorsAreClassifiedWithoutExposingResponse) { + for (const auto status : {HTTP_BAD_REQUEST, HTTP_UNAUTHORIZED, HTTP_REQUEST_TIME_OUT, HTTP_TOO_MANY_REQUESTS, + HTTP_INTERNAL_SERVER_ERROR, HTTP_BAD_GATEWAY, HTTP_SERVICE_UNAVAILABLE, HTTP_GATEWAY_TIME_OUT}) { + for (const TString& body : {"private-response", "{}", R"({"error":42})", R"({"error":"private-error"})", R"({"error":"invalid_client"})"}) { + TOidcTestServer server; + server.Enqueue(body, status); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + try { + protocol.ClientGrant(); + UNIT_FAIL("Expected HTTP error"); + } catch (const TError& error) { + UNIT_ASSERT_VALUES_EQUAL(error.Retryable, status != HTTP_BAD_REQUEST && status != HTTP_UNAUTHORIZED); + UNIT_ASSERT_VALUES_EQUAL(error.Code, body.Contains("invalid_client") ? "invalid_client" : ""); + UNIT_ASSERT(!TString(error.what()).Contains("private")); + } + } + } +} + +Y_UNIT_TEST(RejectsOversizedResponse) { + TOidcTestServer server; + server.Enqueue(TString(1024 * 1024 + 1, 'x'), HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.ClientGrant(), TError, "response exceeds size limit"); +} + +Y_UNIT_TEST(DeviceRequiresAcceptorAndDiscoveryEndpoint) { + TOidcTestServer server; + auto config = DeviceConfig(server); + config.Acceptor(nullptr); + NThreading::TCancellationTokenSource cancellation; + TProtocol withoutAcceptor(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(withoutAcceptor.DeviceGrant([](TDuration) { return true; }), TError, "auth acceptor"); + UNIT_ASSERT_VALUES_EQUAL(server.DiscoveryCount(), 0); + config.Acceptor(std::make_shared()); + auto metadata = Metadata(server); + metadata.EraseValue("device_authorization_endpoint"); + server.SetDiscoveryReply(NJson::WriteJson(metadata, false), HTTP_OK); + TProtocol withoutEndpoint(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(withoutEndpoint.DeviceGrant([](TDuration) { return true; }), TError, "device_authorization_endpoint"); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(DeviceRejectsMissingExpiryAndInvalidVerificationLinks) { + for (const auto& field : {"expires_in", "verification_uri", "verification_uri_complete", "device_code", "user_code"}) { + TOidcTestServer server; + auto response = DeviceResponse(server); + if (TString(field).StartsWith("verification_uri")) { + response[field] = "http://insecure.example"; + } else { + response.EraseValue(field); + } + server.Enqueue(NJson::WriteJson(response, false), HTTP_OK); + const auto config = DeviceConfig(server); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION(protocol.DeviceGrant([](TDuration) { return true; }), std::exception); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); + } +} + +Y_UNIT_TEST(DevicePollingHandlesPendingSlowDownAndTransientErrors) { + TOidcTestServer server; + server.Enqueue(NJson::WriteJson(DeviceResponse(server), false), HTTP_OK); + server.Enqueue(R"({"error":"authorization_pending"})", HTTP_BAD_REQUEST); + server.Enqueue(R"({"error":"slow_down"})", HTTP_BAD_REQUEST); + server.Enqueue(R"({"error":"temporarily_unavailable"})", HTTP_SERVICE_UNAVAILABLE); + server.Enqueue(R"({"access_token":"device-access","token_type":"Bearer","expires_in":600})", HTTP_OK); + const auto config = DeviceConfig(server); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + std::vector delays; + const auto result = protocol.DeviceGrant([&](TDuration delay) { + delays.push_back(delay); + return true; + }); + UNIT_ASSERT_VALUES_EQUAL(result.AccessToken.Token, "device-access"); + UNIT_ASSERT_VALUES_EQUAL(delays.size(), 4); + UNIT_ASSERT_VALUES_EQUAL(delays[0], TDuration::Seconds(5)); + UNIT_ASSERT_VALUES_EQUAL(delays[1], TDuration::Seconds(5)); + UNIT_ASSERT_VALUES_EQUAL(delays[2], TDuration::Seconds(10)); + UNIT_ASSERT_VALUES_EQUAL(delays[3], TDuration::Seconds(20)); + const auto requests = server.Requests(); + UNIT_ASSERT_VALUES_EQUAL(requests.size(), 5); + UNIT_ASSERT_VALUES_EQUAL(requests[0].Form.Get("scope"), "read openid"); + for (size_t i = 1; i < requests.size(); ++i) { + UNIT_ASSERT_VALUES_EQUAL(requests[i].Form.Get("client_id"), "public-client"); + UNIT_ASSERT(requests[i].Authorization.empty()); + } +} + +Y_UNIT_TEST(DeviceDenialIsTerminal) { + TOidcTestServer server; + server.Enqueue(NJson::WriteJson(DeviceResponse(server), false), HTTP_OK); + server.Enqueue(R"({"error":"access_denied"})", HTTP_BAD_REQUEST); + const auto config = DeviceConfig(server); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.DeviceGrant([](TDuration) { return true; }), TError, "access_denied"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 2); +} + +Y_UNIT_TEST(CancellationWaitsForPendingHttpRequest) { + TOidcTestServer server; + auto gate = NThreading::NewPromise(); + server.BlockTokenRepliesUntil(gate.GetFuture()); + server.Enqueue(R"({"access_token":"late","token_type":"Bearer","expires_in":600})", HTTP_OK); + const auto config = server.ClientConfig(); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + auto result = std::async(std::launch::async, [&] { return protocol.ClientGrant(); }); + const bool requested = server.WaitRequests(1); + cancellation.Cancel(); + const bool stillRunning = result.wait_for(std::chrono::milliseconds(100)) == std::future_status::timeout; + gate.TrySetValue(); + UNIT_ASSERT(requested); + UNIT_ASSERT(stillRunning); + UNIT_ASSERT_EXCEPTION(result.get(), std::exception); +} +Y_UNIT_TEST(DiscoveryIssuerComparisonPreservesTrailingSlash) { + TOidcTestServer server; + auto config = server.ClientConfig(); + config.Issuer += "/"; + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.ClientGrant(), TError, "issuer mismatch"); + UNIT_ASSERT(server.Requests().empty()); +} + +Y_UNIT_TEST(DeviceAcceptsOpaqueTokenWithoutLifetime) { + TOidcTestServer server; + server.Enqueue(NJson::WriteJson(DeviceResponse(server), false), HTTP_OK); + server.Enqueue(R"({"access_token":"opaque","token_type":"Bearer","refresh_token":"refresh"})", HTTP_OK); + const auto config = DeviceConfig(server); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + const auto result = protocol.DeviceGrant([](TDuration) { return true; }); + UNIT_ASSERT_VALUES_EQUAL(result.AccessToken.Token, "opaque"); + UNIT_ASSERT(!result.AccessToken.ExpiresAt.has_value()); + UNIT_ASSERT(result.RefreshToken.has_value()); + UNIT_ASSERT_VALUES_EQUAL(result.RefreshToken->Token, "refresh"); +} + +Y_UNIT_TEST(DeviceDoesNotPollAfterWaitOvershootsExpiry) { + TOidcTestServer server; + auto response = DeviceResponse(server); + response["expires_in"] = 1; + server.Enqueue(NJson::WriteJson(response, false), HTTP_OK); + auto acceptor = std::make_shared(); + const auto config = DeviceConfig(server).Acceptor(acceptor); + NThreading::TCancellationTokenSource cancellation; + TProtocol protocol(config, cancellation.Token()); + UNIT_ASSERT_EXCEPTION_CONTAINS(protocol.DeviceGrant([&](TDuration) { + NThreading::NewPromise().GetFuture().Wait(acceptor->Wait().ExpiresAt + TDuration::MilliSeconds(1)); + return true; + }), TError, "device authorization expired"); + UNIT_ASSERT_VALUES_EQUAL(server.Requests().size(), 1); +} + +} diff --git a/tests/unit/client/oidc/test_server.cpp b/tests/unit/client/oidc/test_server.cpp new file mode 100644 index 00000000000..8d105035fc9 --- /dev/null +++ b/tests/unit/client/oidc/test_server.cpp @@ -0,0 +1,70 @@ +#include "test_server.h" + +#include + +#include +#include + +std::optional TMemoryTokenCacher::Read() const { + with_lock (Mutex) { + return Tokens; + } +} + +void TMemoryTokenCacher::Write(const NYdb::NOidc::TTokenCache& tokens) { + with_lock (Mutex) { + Tokens = tokens; + } +} + +void TTestAcceptor::Accept(const NYdb::NOidc::TDeviceAuthInfo& info) { + with_lock (Mutex) { + Info = info; + } + Changed.notify_all(); +} + +NYdb::NOidc::TDeviceAuthInfo TTestAcceptor::Wait() { + with_lock (Mutex) { + UNIT_ASSERT(Changed.wait_for(Mutex, std::chrono::seconds(10), [&] { return Info.has_value(); })); + return Info.value(); + } +} + +void TQueuedOidcFacility::AddPeriodicTask(NYdb::TPeriodicCb&&, NYdb::TDeadline::Duration) { +} + +void TQueuedOidcFacility::PostToResponseQueue(NYdb::TPostTaskCb&& callback) { + with_lock (Mutex) { + Tasks.push_back(std::move(callback)); + } + Changed.notify_all(); +} + +bool TQueuedOidcFacility::WaitForTask() { + with_lock (Mutex) { + return Changed.wait_for(Mutex, std::chrono::seconds(10), [&] { return !Tasks.empty(); }); + } +} + +void TQueuedOidcFacility::RunTasks() { + std::vector tasks; + with_lock (Mutex) { + tasks.swap(Tasks); + } + for (auto& task : tasks) { + task(); + } +} + +void TQueuedOidcFacility::DiscardTasks() { + with_lock (Mutex) { + Tasks.clear(); + } +} + +void TGatedOidcCacher::Write(const NYdb::NOidc::TTokenCache& tokens) { + Entered.TrySetValue(); + Release.GetFuture().Wait(); + TMemoryTokenCacher::Write(tokens); +} diff --git a/tests/unit/client/oidc/test_server.h b/tests/unit/client/oidc/test_server.h new file mode 100644 index 00000000000..b8fa00cc449 --- /dev/null +++ b/tests/unit/client/oidc/test_server.h @@ -0,0 +1,57 @@ +#pragma once + +#include +#include + +#include + +#include + +class TMemoryTokenCacher: public NYdb::NOidc::ITokenCacher { +public: + std::optional Read() const override; + + void Write(const NYdb::NOidc::TTokenCache& tokens) override; + +private: + mutable TMutex Mutex; + std::optional Tokens; +}; + +class TTestAcceptor: public NYdb::NOidc::IAuthAcceptor { +public: + void Accept(const NYdb::NOidc::TDeviceAuthInfo& info) override; + + NYdb::NOidc::TDeviceAuthInfo Wait(); + +private: + TMutex Mutex; + std::condition_variable_any Changed; + std::optional Info; +}; + +class TQueuedOidcFacility: public NYdb::ICoreFacility { +public: + void AddPeriodicTask(NYdb::TPeriodicCb&&, NYdb::TDeadline::Duration) override; + + void PostToResponseQueue(NYdb::TPostTaskCb&& callback) override; + + bool WaitForTask(); + + void RunTasks(); + + void DiscardTasks(); + +private: + TMutex Mutex; + std::condition_variable_any Changed; + std::vector Tasks; +}; + +class TGatedOidcCacher: public TMemoryTokenCacher { +public: + void Write(const NYdb::NOidc::TTokenCache& tokens) override; + + NThreading::TPromise Entered = NThreading::NewPromise(); + NThreading::TPromise Release = NThreading::NewPromise(); +}; diff --git a/tests/unit/client/table/table_ut.cpp b/tests/unit/client/table/table_ut.cpp index 61fc95843a0..78edc65199d 100644 --- a/tests/unit/client/table/table_ut.cpp +++ b/tests/unit/client/table/table_ut.cpp @@ -895,3 +895,25 @@ TEST(TableTest, AlterTableSetMetricsSettings) { Ydb::Table::MetricsSettings::METRICS_LEVEL_PARTITION ); } + +TEST(TtlTierSettings, ObjectKeyPrefixRoundTrip) { + using namespace NYdb::NTable; + for (auto&& prefix : {std::optional(), std::make_optional(std::string()), + std::make_optional(std::string("archive//2026:09/"))}) { + const TTtlTierSettings original(TDateTypeColumnModeSettings("ts", TDuration::Days(1)), + TTtlEvictToExternalStorageAction("/Root/eds", prefix)); + Ydb::Table::TtlTier proto; + original.SerializeTo(proto); + ASSERT_EQ(proto.evict_to_external_storage().has_object_key_prefix(), prefix.has_value()); + const auto restored = TTtlTierSettings::FromProto(proto); + ASSERT_TRUE(restored); + const auto& action = std::get(restored->GetAction()); + EXPECT_EQ(action.GetStorage(), "/Root/eds"); + EXPECT_EQ(action.GetObjectKeyPrefix(), prefix); + } + + Ydb::Table::EvictionToExternalStorageSettings proto; + proto.set_object_key_prefix("previous"); + TTtlEvictToExternalStorageAction("/Root/eds").SerializeTo(proto); + EXPECT_FALSE(proto.has_object_key_prefix()); +} diff --git a/tests/unit/client/topic/CMakeLists.txt b/tests/unit/client/topic/CMakeLists.txt index 769ed123c57..e59cab4371e 100644 --- a/tests/unit/client/topic/CMakeLists.txt +++ b/tests/unit/client/topic/CMakeLists.txt @@ -33,3 +33,16 @@ add_ydb_test(NAME client-topic-read-session-credentials_ut LABELS unit ) + +add_ydb_test(NAME client-topic-read-session-accounting_ut + INCLUDE_DIRS ${YDB_SDK_SOURCE_DIR} + SOURCES ${YDB_SDK_SOURCE_DIR}/src/client/topic/ut/read_session_accounting_ut.cpp + LINK_LIBRARIES + YDB-CPP-SDK::Topic + client-ydb_topic-impl + LABELS unit +) + +# read_session_kafka_timestamps_ut.cpp and reset_offset_ut.cpp use the +# upstream TTopicSdkTestSetup and ydb/core server fixtures, which are absent +# from the standalone SDK repository. diff --git a/tests/unit/client/value/embedding_ut.cpp b/tests/unit/client/value/embedding_ut.cpp new file mode 100644 index 00000000000..a705740ea82 --- /dev/null +++ b/tests/unit/client/value/embedding_ut.cpp @@ -0,0 +1,55 @@ +#include + +#include + +#include + +#include +#include +#include +#include + +namespace NYdb { + +TEST(Embedding, Float32) { + const TVector values = {1.0f, -2.0f, 0.5f}; + const auto value = NValueHelpers::Embedding(values); + TValueParser parser(value); + + EXPECT_EQ(parser.GetPrimitiveType(), EPrimitiveType::Bytes); + EXPECT_EQ(parser.GetBytes(), std::string("\x00\x00\x80\x3f\x00\x00\x00\xc0\x00\x00\x00\x3f\x01", 13)); +} + +TEST(Embedding, IntegersAreConvertedToFloat32) { + const std::array values = {-2, 16777217}; + const auto value = NValueHelpers::Embedding(values); + TValueParser parser(value); + + EXPECT_EQ(parser.GetBytes(), std::string("\x00\x00\x00\xc0\x00\x00\x80\x4b\x01", 9)); +} + +TEST(Embedding, UnsignedIntegersAreConvertedToFloat32) { + const std::array values = {2}; + const auto value = NValueHelpers::Embedding(values); + TValueParser parser(value); + + EXPECT_EQ(parser.GetBytes(), std::string("\x00\x00\x00\x40\x01", 5)); +} + +TEST(Embedding, Float64IsConvertedToFloat32) { + const std::vector values = {1.00000001}; + const auto value = NValueHelpers::Embedding(values); + TValueParser parser(value); + + EXPECT_EQ(parser.GetBytes(), std::string("\x00\x00\x80\x3f\x01", 5)); +} + +TEST(Embedding, EmptyHasFormatByte) { + const auto value = NValueHelpers::Embedding(std::vector{}); + TValueParser parser(value); + + EXPECT_EQ(parser.GetPrimitiveType(), EPrimitiveType::Bytes); + EXPECT_EQ(parser.GetBytes(), std::string("\x01", 1)); +} + +} // namespace NYdb diff --git a/tests/unit/library/grpc_client/grpc_client_low_ut.cpp b/tests/unit/library/grpc_client/grpc_client_low_ut.cpp index a53ee718334..171473c3a84 100644 --- a/tests/unit/library/grpc_client/grpc_client_low_ut.cpp +++ b/tests/unit/library/grpc_client/grpc_client_low_ut.cpp @@ -12,6 +12,111 @@ class TTestStub { {} }; +namespace NYdbGrpc::inline V3 { + +struct TStreamRequestReadWriteProcessorTestAccess { + template + static void SetStream(TProcessor& processor, typename TProcessor::TAsyncReaderWriterPtr stream) { + processor.Stream = std::move(stream); + processor.Started = true; + processor.ConnectedCallback = nullptr; + } +}; + +} // namespace NYdbGrpc::inline V3 + +namespace { + +struct TTestMessage { + int Value = 0; + + void Swap(TTestMessage* other) { + std::swap(Value, other->Value); + } +}; + +class TTestAsyncReaderWriter final + : public grpc::ClientAsyncReaderWriterInterface +{ +public: + void StartCall(void*) override { + } + + void ReadInitialMetadata(void*) override { + } + + void Finish(grpc::Status* status, void* tag) override { + FinishStatus = status; + FinishTag = tag; + } + + void Write(const TTestMessage& message, void* tag) override { + WrittenValues.push_back(message.Value); + WriteTag = tag; + } + + void Write(const TTestMessage&, grpc::WriteOptions, void*) override { + Y_ABORT("Unexpected Write with options"); + } + + void Read(TTestMessage*, void*) override { + Y_ABORT("Unexpected Read"); + } + + void WritesDone(void* tag) override { + WritesDoneTag = tag; + ++WritesDoneCalls; + } + + void CompleteWrite(bool ok) { + Complete(WriteTag, ok); + } + + void CompleteWritesDone(bool ok) { + Complete(WritesDoneTag, ok); + } + + void CompleteFinish(const grpc::Status& status, bool ok = true) { + UNIT_ASSERT(FinishStatus); + *FinishStatus = status; + Complete(FinishTag, ok); + } + + std::vector WrittenValues; + size_t WritesDoneCalls = 0; + +private: + static void Complete(void*& tag, bool ok) { + UNIT_ASSERT(tag); + auto* event = static_cast(std::exchange(tag, nullptr)); + event->Execute(ok); + event->Destroy(); + } + +private: + grpc::Status* FinishStatus = nullptr; + void* WriteTag = nullptr; + void* WritesDoneTag = nullptr; + void* FinishTag = nullptr; +}; + +using TTestProcessor = TStreamRequestReadWriteProcessor; + +struct TProcessorFixture { + TProcessorFixture() + : Processor(MakeIntrusive([](TGrpcStatus&&, auto) {})) + { + auto stream = std::make_unique(); + Stream = stream.get(); + TStreamRequestReadWriteProcessorTestAccess::SetStream(*Processor, std::move(stream)); + } + + TIntrusivePtr Processor; + TTestAsyncReaderWriter* Stream = nullptr; +}; + +} // anonymous namespace + Y_UNIT_TEST_SUITE(ChannelPoolTests) { Y_UNIT_TEST(UnusedStubsHoldersDeletion) { TGRpcClientConfig clientConfig("invalid_host:invalid_port"); @@ -59,3 +164,80 @@ Y_UNIT_TEST_SUITE(ChannelPoolTests) { } } // ChannelPoolTests ut suite + +Y_UNIT_TEST_SUITE(StreamRequestReadWriteProcessorTests) { + Y_UNIT_TEST(RejectsOperationsQueuedAfterWritesDone) { + TProcessorFixture fixture; + std::vector firstWriteStatuses; + std::vector writesDoneStatuses; + std::vector lateWriteStatuses; + std::vector repeatedWritesDoneStatuses; + + fixture.Processor->Write(TTestMessage{1}, [&](TGrpcStatus&& status) { + firstWriteStatuses.push_back(std::move(status)); + }); + fixture.Processor->WritesDone([&](TGrpcStatus&& status) { + writesDoneStatuses.push_back(std::move(status)); + }); + fixture.Processor->Write(TTestMessage{2}, [&](TGrpcStatus&& status) { + lateWriteStatuses.push_back(std::move(status)); + }); + fixture.Processor->WritesDone([&](TGrpcStatus&& status) { + repeatedWritesDoneStatuses.push_back(std::move(status)); + }); + + UNIT_ASSERT_VALUES_EQUAL(fixture.Stream->WrittenValues.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(fixture.Stream->WrittenValues.front(), 1); + UNIT_ASSERT_VALUES_EQUAL(fixture.Stream->WritesDoneCalls, 0); + UNIT_ASSERT_VALUES_EQUAL(lateWriteStatuses.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(lateWriteStatuses.front().GRpcStatusCode, static_cast(grpc::StatusCode::FAILED_PRECONDITION)); + UNIT_ASSERT_VALUES_EQUAL(repeatedWritesDoneStatuses.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(repeatedWritesDoneStatuses.front().GRpcStatusCode, static_cast(grpc::StatusCode::FAILED_PRECONDITION)); + + fixture.Stream->CompleteWrite(true); + UNIT_ASSERT_VALUES_EQUAL(firstWriteStatuses.size(), 1); + UNIT_ASSERT(firstWriteStatuses.front().Ok()); + UNIT_ASSERT_VALUES_EQUAL(fixture.Stream->WritesDoneCalls, 1); + + fixture.Stream->CompleteWritesDone(true); + UNIT_ASSERT_VALUES_EQUAL(writesDoneStatuses.size(), 1); + UNIT_ASSERT(writesDoneStatuses.front().Ok()); + } + + Y_UNIT_TEST(DeliversWritesDoneFailureAfterFinish) { + TProcessorFixture fixture; + std::vector writesDoneStatuses; + std::vector finishStatuses; + + fixture.Processor->WritesDone([&](TGrpcStatus&& status) { + writesDoneStatuses.push_back(std::move(status)); + }); + fixture.Stream->CompleteWritesDone(false); + UNIT_ASSERT(writesDoneStatuses.empty()); + + fixture.Processor->Finish([&](TGrpcStatus&& status) { + finishStatuses.push_back(std::move(status)); + }); + fixture.Stream->CompleteFinish(grpc::Status(grpc::StatusCode::INTERNAL, "half-close failed")); + + UNIT_ASSERT_VALUES_EQUAL(writesDoneStatuses.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(writesDoneStatuses.front().GRpcStatusCode, static_cast(grpc::StatusCode::INTERNAL)); + UNIT_ASSERT_VALUES_EQUAL(finishStatuses.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(finishStatuses.front().GRpcStatusCode, static_cast(grpc::StatusCode::INTERNAL)); + } + + Y_UNIT_TEST(CancelWhileWritesDoneIsInFlight) { + TProcessorFixture fixture; + std::vector writesDoneStatuses; + + fixture.Processor->WritesDone([&](TGrpcStatus&& status) { + writesDoneStatuses.push_back(std::move(status)); + }); + fixture.Processor->Cancel(); + fixture.Stream->CompleteWritesDone(false); + fixture.Stream->CompleteFinish(grpc::Status(grpc::StatusCode::CANCELLED, "cancelled")); + + UNIT_ASSERT_VALUES_EQUAL(writesDoneStatuses.size(), 1); + UNIT_ASSERT_VALUES_EQUAL(writesDoneStatuses.front().GRpcStatusCode, static_cast(grpc::StatusCode::CANCELLED)); + } +} // StreamRequestReadWriteProcessorTests suite diff --git a/util/CMakeLists.txt b/util/CMakeLists.txt index 2da80278392..35fcc0330c4 100644 --- a/util/CMakeLists.txt +++ b/util/CMakeLists.txt @@ -726,13 +726,8 @@ if (NOT WIN32) endif () if (CMAKE_SYSTEM_PROCESSOR STREQUAL "x86_64") - target_yasm_source(yutil - PRIVATE - ${YDB_SDK_SOURCE_DIR}/util/system/context_x86.asm - -I - ${YDB_SDK_BINARY_DIR} - -I - ${YDB_SDK_SOURCE_DIR} + target_sources(yutil PRIVATE + ${YDB_SDK_SOURCE_DIR}/util/system/context_x86.S ) elseif (CMAKE_SYSTEM_PROCESSOR STREQUAL "aarch64" OR CMAKE_SYSTEM_PROCESSOR STREQUAL "arm64") target_sources(yutil PRIVATE diff --git a/util/charset/wide_ut.cpp b/util/charset/wide_ut.cpp index e221ec4ad66..5192188a883 100644 --- a/util/charset/wide_ut.cpp +++ b/util/charset/wide_ut.cpp @@ -671,7 +671,9 @@ class TWideUtilTest: public TTestBase { Collapse(s); UNIT_ASSERT(s == w); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == w.c_str()); // Collapse() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == w.c_str()); // Collapse() does not change the string at all + } #endif } s = ASCIIToWide(" 123 456 "); @@ -699,7 +701,9 @@ class TWideUtilTest: public TTestBase { Collapse(s); UNIT_ASSERT(s == w); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == w.c_str()); // Collapse() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == w.c_str()); // Collapse() does not change the string at all + } #endif } s = ASCIIToWide(" "); @@ -832,19 +836,25 @@ class TWideUtilTest: public TTestBase { Strip(s); UNIT_ASSERT(s == w); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + } #endif s = w; StripLeft(s); UNIT_ASSERT(s == w); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + } #endif s = w; StripRight(s); UNIT_ASSERT(s == w); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == w.c_str()); // Strip() does not change the string at all + } #endif } @@ -1157,7 +1167,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToLower(writableCopy.Detach(), writableCopy.size())); @@ -1178,7 +1190,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToLower(writableCopy.Detach(), writableCopy.size())); @@ -1198,7 +1212,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s, 100500)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToLowerRet(copy, 100500) == lower); @@ -1212,7 +1228,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s, 100500, 1111)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToLowerRet(copy, 100500, 1111) == lower); @@ -1245,7 +1263,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToLower(writableCopy.Detach(), writableCopy.size())); @@ -1266,7 +1286,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToLower(writableCopy.Detach(), writableCopy.size())); @@ -1315,7 +1337,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s, 2)); UNIT_ASSERT(s == lower); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToLowerRet(copy, 2) == lower); @@ -1340,7 +1364,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s, 3, 1)); UNIT_ASSERT(s == copy); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToLowerRet(copy, 3, 1) == lower); @@ -1354,7 +1380,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToLower(s, 3, 100500)); UNIT_ASSERT(s == copy); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToLowerRet(copy, 3, 100500) == lower); @@ -1372,7 +1400,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToUpper(s)); UNIT_ASSERT(s == upper); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToUpper(writableCopy.Detach(), writableCopy.size())); @@ -1393,7 +1423,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToUpper(s)); UNIT_ASSERT(s == upper); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToUpper(writableCopy.Detach(), writableCopy.size())); @@ -1414,7 +1446,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToUpper(s, 100500)); UNIT_ASSERT(s == upper); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToUpper(writableCopy.Detach(), writableCopy.size())); @@ -1434,7 +1468,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToUpper(s, 100500, 1111)); UNIT_ASSERT(s == upper); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToUpperRet(copy, 100500, 1111) == upper); @@ -1467,7 +1503,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToUpper(s)); UNIT_ASSERT(s == copy); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToUpper(writableCopy.Detach(), writableCopy.size())); @@ -1589,7 +1627,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToTitle(writableCopy.Detach(), writableCopy.size())); @@ -1610,7 +1650,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToTitle(writableCopy.Detach(), writableCopy.size())); @@ -1630,7 +1672,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s, 100500)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToTitleRet(copy) == title); @@ -1644,7 +1688,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s, 100500, 1111)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToTitleRet(copy) == title); @@ -1677,7 +1723,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToTitle(writableCopy.Detach(), writableCopy.size())); @@ -1716,7 +1764,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(!ToTitle(writableCopy.Detach(), writableCopy.size())); @@ -1765,7 +1815,9 @@ class TWideUtilTest: public TTestBase { UNIT_ASSERT(!ToTitle(s, 2)); UNIT_ASSERT(s == title); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.data() == copy.data()); + if (TStringUseCow) { + UNIT_ASSERT(s.data() == copy.data()); + } #endif UNIT_ASSERT(ToTitleRet(copy, 2) == title); diff --git a/util/folder/path_ut.cpp b/util/folder/path_ut.cpp index 095969f3bb4..eb13d479ee3 100644 --- a/util/folder/path_ut.cpp +++ b/util/folder/path_ut.cpp @@ -705,11 +705,12 @@ Y_UNIT_TEST_SUITE(TFsPathTests) { // Broken since f8699c0a71a528d287b84cd0bc5b5bb7cec924f0 (5.11 wine-version) return; } - Chmod(testSubdir.c_str(), 0); - Y_DEFER { - Chmod(testSubdir.c_str(), MODE0777); - }; - TWinFileDenyAccessScope dirAcl(testDir, FILE_WRITE_DATA); + // Windows removes a directory for whoever holds DELETE on it or + // FILE_DELETE_CHILD on its parent, so both have to be taken away. The + // read-only attribute is not one of the ways to stop it: NFs::Remove + // drops that attribute first, on purpose, to match unix. + TWinFileDenyAccessScope subdirAcl(testSubdir, DELETE); + TWinFileDenyAccessScope dirAcl(testDir, FILE_DELETE_CHILD); #else Chmod(testDir.c_str(), S_IRUSR | S_IXUSR); Y_DEFER { diff --git a/util/generic/fastqueue.h b/util/generic/fastqueue.h index 1fee5b86f6e..29b0bcb8b69 100644 --- a/util/generic/fastqueue.h +++ b/util/generic/fastqueue.h @@ -35,6 +35,12 @@ class TFastQueue { return tmp->Obj; } + inline const T& Front() const noexcept { + Y_ASSERT(!this->Empty()); + + return Queue_.Back()->Obj; + } + inline size_t Size() const noexcept { return Size_; } diff --git a/util/generic/guid.cpp b/util/generic/guid.cpp index 0d20c467ccb..fdb956a1e7c 100644 --- a/util/generic/guid.cpp +++ b/util/generic/guid.cpp @@ -6,9 +6,16 @@ #include #include #include +#include #include namespace { + constexpr ui32 UUID_V7_VERSION_MASK = 0x0000f000; + constexpr ui32 UUID_V7_VERSION = 0x00007000; + constexpr ui32 UUID_VARIANT_MASK = 0xc0000000; + constexpr ui32 UUID_RFC_VARIANT = 0x80000000; + constexpr ui64 UUID_V7_TIMESTAMP_MASK = (ui64{1} << 48) - 1; + inline void LowerCaseHex(TString& s) { for (auto&& c : s) { c = AsciiToLower(c); @@ -65,6 +72,31 @@ TGUID TGUID::CreateTimebased() { return result; } +TGUID TGUID::CreateUuidV7() { + const ui64 timestamp = MilliSeconds() & UUID_V7_TIMESTAMP_MASK; + const ui64 randB = RandomNumber(); + + TGUID result; + result.dw[0] = static_cast(timestamp >> 16); + result.dw[1] = static_cast(timestamp << 16) | + UUID_V7_VERSION | + (RandomNumber() & 0x0fff); + result.dw[2] = (static_cast(randB >> 32) & ~UUID_VARIANT_MASK) | UUID_RFC_VARIANT; + result.dw[3] = static_cast(randB); + return result; +} + +bool IsUuidV7(const TGUID& uuid) noexcept { + return (uuid.dw[1] & UUID_V7_VERSION_MASK) == UUID_V7_VERSION && + (uuid.dw[2] & UUID_VARIANT_MASK) == UUID_RFC_VARIANT; +} + +TInstant GetUuidV7Timestamp(const TGUID& uuid) noexcept { + Y_ASSERT(IsUuidV7(uuid)); + const ui64 timestamp = (static_cast(uuid.dw[0]) << 16) | (uuid.dw[1] >> 16); + return TInstant::MilliSeconds(timestamp); +} + TString GetGuidAsString(const TGUID& g) { return g.AsGuidString(); } diff --git a/util/generic/guid.h b/util/generic/guid.h index 93e06a81f87..ab26ed48e83 100644 --- a/util/generic/guid.h +++ b/util/generic/guid.h @@ -2,6 +2,7 @@ #include "fwd.h" +#include #include /** @@ -42,6 +43,11 @@ struct TGUID { * https://datatracker.ietf.org/doc/html/rfc4122#section-4.1 **/ static TGUID CreateTimebased(); + + /** + * Generate a time-ordered UUID version 7 as specified by RFC 9562. + **/ + static TGUID CreateUuidV7(); }; constexpr bool operator==(const TGUID& a, const TGUID& b) noexcept { @@ -77,3 +83,14 @@ bool GetGuid(TStringBuf s, TGUID& result); **/ TGUID GetUuid(TStringBuf s); bool GetUuid(TStringBuf s, TGUID& result); + +/** + * Returns true if uuid is an RFC 9562 UUID version 7 with the RFC variant. + **/ +bool IsUuidV7(const TGUID& uuid) noexcept; + +/** + * Extracts the timestamp from a UUID version 7. + * Asserts that uuid is a valid UUIDv7. + **/ +TInstant GetUuidV7Timestamp(const TGUID& uuid) noexcept; diff --git a/util/generic/guid_ut.cpp b/util/generic/guid_ut.cpp index f6a155b89e3..7d912d5a0b5 100644 --- a/util/generic/guid_ut.cpp +++ b/util/generic/guid_ut.cpp @@ -2,6 +2,10 @@ #include "guid.h" +#include + +#include + Y_UNIT_TEST_SUITE(TGuidTest) { // TODO - make real constructor static TGUID Construct(ui32 d1, ui32 d2, ui32 d3, ui32 d4) { @@ -125,4 +129,42 @@ Y_UNIT_TEST_SUITE(TGuidTest) { UNIT_ASSERT(!guid.empty()); UNIT_ASSERT_EQUAL(guid[14], '1'); } + + Y_UNIT_TEST(UuidV7) { + // UUIDv7 stores timestamps with millisecond precision, while TInstant::Now() + // has microsecond precision. Round the bounds to avoid rejecting a valid UUID + // whose extracted timestamp is earlier within the same millisecond. + const TInstant before = TInstant::MilliSeconds(TInstant::Now().MilliSeconds()); + const TGUID uuid = TGUID::CreateUuidV7(); + const TInstant after = TInstant::MilliSeconds(TInstant::Now().MilliSeconds()); + + UNIT_ASSERT(IsUuidV7(uuid)); + UNIT_ASSERT_GE(GetUuidV7Timestamp(uuid), before); + UNIT_ASSERT_LE(GetUuidV7Timestamp(uuid), after); + UNIT_ASSERT_EQUAL(uuid.AsUuidString()[14], '7'); + } + + Y_UNIT_TEST(UuidV7RfcExample) { + // RFC 9562, Appendix A.6. + const TGUID uuid = GetUuid("017f22e2-79b0-7cc3-98c4-dc0c0c07398f"); + + UNIT_ASSERT(IsUuidV7(uuid)); + UNIT_ASSERT_EQUAL(GetUuidV7Timestamp(uuid), TInstant::MilliSeconds(1645557742000ULL)); + } + + Y_UNIT_TEST(UuidV7RejectsWrongVersionAndVariant) { + UNIT_ASSERT(!IsUuidV7(GetUuid("017f22e2-79b0-4cc3-98c4-dc0c0c07398f"))); + UNIT_ASSERT(!IsUuidV7(GetUuid("017f22e2-79b0-7cc3-18c4-dc0c0c07398f"))); + } + + Y_UNIT_TEST(UuidV7ValuesAreDistinct) { + std::unordered_set uuids; + constexpr size_t count = 1000; + + for (size_t i = 0; i < count; ++i) { + uuids.insert(TGUID::CreateUuidV7().AsUuidString()); + } + + UNIT_ASSERT_EQUAL(uuids.size(), count); + } } // Y_UNIT_TEST_SUITE(TGuidTest) diff --git a/util/generic/ptr.h b/util/generic/ptr.h index 7219a1c85ca..9f6b5c65273 100644 --- a/util/generic/ptr.h +++ b/util/generic/ptr.h @@ -7,6 +7,7 @@ #include "typetraits.h" #include "singleton.h" +#include #include #include @@ -259,6 +260,7 @@ class TAutoPtr: public TPointerBase, T> { mutable T* T_; }; +// Deprecated, use std::unique_ptr instead template class Y_TRIVIAL_ABI THolder: public TPointerBase, T> { public: @@ -299,6 +301,13 @@ class Y_TRIVIAL_ABI THolder: public TPointerBase, T> { { } + template && + !std::is_array_v && std::is_convertible_v>> + explicit THolder(std::unique_ptr&& that) noexcept + : T_(that.release()) + { + } + THolder(const THolder&) = delete; THolder& operator=(const THolder&) = delete; @@ -344,6 +353,18 @@ class Y_TRIVIAL_ABI THolder: public TPointerBase, T> { return T_; } + inline T* release() noexcept Y_WARN_UNUSED_RESULT { + return Release(); + } + + Y_REINITIALIZES_OBJECT inline void reset(T* t) noexcept { + Reset(t); + } + + Y_REINITIALIZES_OBJECT inline void reset() noexcept { + Reset(); + } + inline operator TAutoPtr() noexcept { return Release(); } @@ -358,6 +379,19 @@ class Y_TRIVIAL_ABI THolder: public TPointerBase, T> { return *this; } + template && + !std::is_array_v && std::is_convertible_v>> + THolder& operator=(std::unique_ptr&& that) noexcept { + this->Reset(that.release()); + return *this; + } + + template && + !std::is_array_v && std::is_convertible_v>> + explicit operator std::unique_ptr() && noexcept { + return std::unique_ptr(Release()); + } + template THolder& operator=(THolder&& that) noexcept { this->Reset(that.Release()); diff --git a/util/generic/ptr.pxd b/util/generic/ptr.pxd index e4078295190..ab5b839650a 100644 --- a/util/generic/ptr.pxd +++ b/util/generic/ptr.pxd @@ -2,10 +2,14 @@ cdef extern from "" nogil: cdef cppclass THolder[T]: THolder(...) T* Get() + T* get() void Destroy() T* Release() + T* release() void Reset() void Reset(T*) + void reset() + void reset(T*) void Swap(THolder[T]) diff --git a/util/generic/ptr_ut.cpp b/util/generic/ptr_ut.cpp index e6a4a7837c0..25390efbc46 100644 --- a/util/generic/ptr_ut.cpp +++ b/util/generic/ptr_ut.cpp @@ -189,6 +189,168 @@ void TPointerTest::TestHolderPtrMoveAssignmentInheritance() { basePtr = THolder(new TDerived); } +Y_UNIT_TEST_SUITE(THolderUniquePtrTest) { + Y_UNIT_TEST(ConstructionAndAssignment) { + UNIT_ASSERT_VALUES_EQUAL(cnt, 0); + { + auto source = std::make_unique(); + auto* raw = source.get(); + THolder holder{std::move(source)}; + UNIT_ASSERT(!source); + UNIT_ASSERT(holder.Get() == raw); + + std::unique_ptr destination = static_cast>(std::move(holder)); + UNIT_ASSERT(!holder); + UNIT_ASSERT(destination.get() == raw); + + holder = MakeHolder(); + UNIT_ASSERT_VALUES_EQUAL(cnt, 2); + UNIT_ASSERT(&(holder = std::move(destination)) == &holder); + UNIT_ASSERT(!destination); + UNIT_ASSERT(holder.Get() == raw); + UNIT_ASSERT_VALUES_EQUAL(cnt, 1); + + destination = std::make_unique(); + UNIT_ASSERT_VALUES_EQUAL(cnt, 2); + destination = static_cast>(std::move(holder)); + UNIT_ASSERT(!holder); + UNIT_ASSERT(destination.get() == raw); + UNIT_ASSERT_VALUES_EQUAL(cnt, 1); + } + UNIT_ASSERT_VALUES_EQUAL(cnt, 0); + } + + Y_UNIT_TEST(EmptyPointers) { + std::unique_ptr source; + THolder holder{std::move(source)}; + UNIT_ASSERT(!holder); + std::unique_ptr destination = static_cast>(std::move(holder)); + UNIT_ASSERT(!destination); + + holder = MakeHolder(42); + holder = std::move(source); + UNIT_ASSERT(!holder); + destination = std::make_unique(42); + destination = static_cast>(std::move(holder)); + UNIT_ASSERT(!destination); + } + + Y_UNIT_TEST(InheritanceAndConst) { + auto source = std::make_unique(); + auto* raw = source.get(); + THolder base{std::move(source)}; + UNIT_ASSERT(!source); + UNIT_ASSERT(base.Get() == raw); + + source = std::make_unique(); + raw = source.get(); + base = std::move(source); + UNIT_ASSERT(!source); + UNIT_ASSERT(base.Get() == raw); + + auto derived = MakeHolder(); + raw = derived.Get(); + std::unique_ptr destination = static_cast>(std::move(derived)); + UNIT_ASSERT(!derived); + UNIT_ASSERT(destination.get() == raw); + + derived = MakeHolder(); + raw = derived.Get(); + destination = static_cast>(std::move(derived)); + UNIT_ASSERT(!derived); + UNIT_ASSERT(destination.get() == raw); + + THolder constHolder{std::make_unique(42)}; + UNIT_ASSERT_VALUES_EQUAL(*constHolder, 42); + constHolder = std::make_unique(43); + UNIT_ASSERT_VALUES_EQUAL(*constHolder, 43); + std::unique_ptr constPtr = static_cast>(MakeHolder(44)); + UNIT_ASSERT_VALUES_EQUAL(*constPtr, 44); + constPtr = static_cast>(MakeHolder(45)); + UNIT_ASSERT_VALUES_EQUAL(*constPtr, 45); + } + + Y_UNIT_TEST(FunctionArguments) { + auto acceptHolder = [](THolder ptr) { return *ptr; }; + auto acceptUnique = [](std::unique_ptr ptr) { return *ptr; }; + // both directions require an explicit cast + UNIT_ASSERT_VALUES_EQUAL(acceptHolder(THolder{std::make_unique(42)}), 42); + UNIT_ASSERT_VALUES_EQUAL(acceptUnique(static_cast>(MakeHolder(43))), 43); + } + + Y_UNIT_TEST(StdStyleMethods) { + { + THolder holder = MakeHolder(42); + UNIT_ASSERT_VALUES_EQUAL(*holder.get(), 42); + + int* raw = holder.release(); + UNIT_ASSERT_VALUES_EQUAL(*raw, 42); + UNIT_ASSERT(!holder); + delete raw; + + holder.reset(new int(43)); + UNIT_ASSERT_VALUES_EQUAL(*holder, 43); + + holder.reset(); + UNIT_ASSERT(!holder); + + holder.reset(new int(44)); + holder.reset(nullptr); + UNIT_ASSERT(!holder); + } + { + // reset() must destroy the old object + auto destroyed = 0; + struct TCounter { + int* Count_; + TCounter(int* count) + : Count_(count) + { + } + ~TCounter() { + ++*Count_; + } + }; + THolder holder = MakeHolder(&destroyed); + holder.reset(new TCounter(&destroyed)); + UNIT_ASSERT_VALUES_EQUAL(destroyed, 1); + holder.reset(); + UNIT_ASSERT_VALUES_EQUAL(destroyed, 2); + } + } + + Y_UNIT_TEST(ConversionConstraints) { + static_assert(std::is_nothrow_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, THolder>); + static_assert(std::is_nothrow_assignable_v&, std::unique_ptr&&>); + static_assert(std::is_nothrow_constructible_v, THolder&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + // assignment is not available: it would require an implicit conversion + static_assert(!std::is_assignable_v&, THolder&&>); + static_assert(!std::is_constructible_v, std::unique_ptr&>); + static_assert(!std::is_assignable_v&, std::unique_ptr&>); + static_assert(!std::is_constructible_v, THolder&>); + static_assert(!std::is_assignable_v&, THolder&>); + static_assert(!std::is_constructible_v, const std::unique_ptr&&>); + static_assert(!std::is_constructible_v, const THolder&&>); + static_assert(!std::is_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_assignable_v&, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + static_assert(!std::is_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + static_assert(!std::is_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + static_assert(!std::is_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + static_assert(!std::is_constructible_v, std::unique_ptr&&>); + static_assert(!std::is_convertible_v&&, std::unique_ptr>); + using TCustomUnique = std::unique_ptr; + static_assert(!std::is_constructible_v, TCustomUnique&&>); + static_assert(!std::is_assignable_v&, TCustomUnique&&>); + static_assert(!std::is_convertible_v&&, TCustomUnique>); + } +} // Y_UNIT_TEST_SUITE(THolderUniquePtrTest) + void TPointerTest::TestMakeHolder() { { auto ptr = MakeHolder(5); diff --git a/util/generic/scope.h b/util/generic/scope.h index 2ae3feb7178..fad9f420621 100644 --- a/util/generic/scope.h +++ b/util/generic/scope.h @@ -3,6 +3,8 @@ #include #include +#include +#include #include namespace NPrivate { @@ -56,3 +58,55 @@ namespace NPrivate { // ok = true; // \endcode #define Y_DEFER Y_SCOPE_EXIT(&) + +// A RAII scope guard. +// By default, invokes the provided callback when destroyed. +// The callback can be invoked earlier by `CallNow()` or cancelled by `Drop()`, +// in these cases the destructor becomes no-op. +template +class TDeferredOnceFunction { + static_assert(std::is_nothrow_invocable_v); + +public: + TDeferredOnceFunction(const F& function) + : Function_{function} + { + } + + TDeferredOnceFunction(F&& function) + : Function_{std::move(function)} + { + } + + ~TDeferredOnceFunction() { + if (Function_) { + (*Function_)(); + } + } + + TDeferredOnceFunction(const TDeferredOnceFunction&) = delete; + TDeferredOnceFunction& operator=(const TDeferredOnceFunction&) = delete; + + TDeferredOnceFunction(TDeferredOnceFunction&& other) { + // Cannot use Function_.swap because F may be not move-assignable (e.g. a lambda with captures) + if (other.Function_.has_value()) { + Function_.emplace(std::move(*other.Function_)); + other.Function_.reset(); + } + } + + // Not sure if there are any valid usecases for assignment. + TDeferredOnceFunction& operator=(TDeferredOnceFunction&&) = delete; + + void CallNow() && noexcept { + (*Function_)(); + Function_.reset(); + } + + void Drop() && noexcept { + Function_.reset(); + } + +private: + std::optional Function_; // Not TMaybe, because scope.h is used in builds with disabled exceptions +}; diff --git a/util/generic/scope_ut.cpp b/util/generic/scope_ut.cpp index 2df66fda576..1f291cab6ab 100644 --- a/util/generic/scope_ut.cpp +++ b/util/generic/scope_ut.cpp @@ -41,7 +41,29 @@ Y_UNIT_TEST_SUITE(ScopeToolsTest) { i = 20; }; } - UNIT_ASSERT_VALUES_EQUAL(i, 20); } + + Y_UNIT_TEST(TestDeferred) { + int i = 0; + + { + TDeferredOnceFunction writeI([&]() noexcept { i = 1; }); + { + auto doWriteI = std::move(writeI); + UNIT_ASSERT_VALUES_EQUAL(i, 0); + } // doWriteI called + UNIT_ASSERT_VALUES_EQUAL(i, 1); + + TDeferredOnceFunction updateI([&]() noexcept { i = 2; }); + UNIT_ASSERT_VALUES_EQUAL(i, 1); + std::move(updateI).CallNow(); + UNIT_ASSERT_VALUES_EQUAL(i, 2); + } // moved-from writeI is no-op + UNIT_ASSERT_VALUES_EQUAL(i, 2); + + TDeferredOnceFunction droppedUpdateI([&]() noexcept { i = 3; }); + std::move(droppedUpdateI).Drop(); + UNIT_ASSERT_VALUES_EQUAL(i, 2); + } } // Y_UNIT_TEST_SUITE(ScopeToolsTest) diff --git a/util/generic/string_ut.h b/util/generic/string_ut.h index 32c179a47e3..bb1798f3bc5 100644 --- a/util/generic/string_ut.h +++ b/util/generic/string_ut.h @@ -553,7 +553,9 @@ class TStringTestImpl { TStringType s7(s6); UNIT_ASSERT(s7 == s6); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s7.c_str() == s6.c_str()); + if (TStringUseCow) { + UNIT_ASSERT(s7.c_str() == s6.c_str()); + } #endif TStringType s8(s7, 1, 3); @@ -648,6 +650,10 @@ class TStringTestImpl { #ifndef TSTRING_IS_STD_STRING void TestRefCount() { + if (!TStringUseCow) { + return; // exercises copy-on-write internals (RefCount/IsDetached) + } + using TStr = TStringType; struct TestStroka: public TStr { @@ -1003,6 +1009,10 @@ class TStringTestImpl { #ifndef TSTRING_IS_STD_STRING void TestCharRef() { + if (!TStringUseCow) { + return; // exercises copy-on-write internals (RefCount/IsDetached) + } + const char_type abc[] = {'a', 'b', 'c', 0}; const char_type bbc[] = {'b', 'b', 'c', 0}; const char_type cbc[] = {'c', 'b', 'c', 0}; diff --git a/util/string/strip.h b/util/string/strip.h index c50cb3023ea..2c2c58d8555 100644 --- a/util/string/strip.h +++ b/util/string/strip.h @@ -96,7 +96,9 @@ struct TStripImpl { return true; } - to = from; + if (&to != &from) { + to = from; + } return false; } diff --git a/util/string/strip_ut.cpp b/util/string/strip_ut.cpp index 7236ca2ebf8..4454c750828 100644 --- a/util/string/strip_ut.cpp +++ b/util/string/strip_ut.cpp @@ -200,7 +200,9 @@ Y_UNIT_TEST_SUITE(TStripStringTest) { UNIT_ASSERT(s == s2); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + } #endif } @@ -217,7 +219,9 @@ Y_UNIT_TEST_SUITE(TStripStringTest) { UNIT_ASSERT(s == s2); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + } #endif } @@ -234,7 +238,9 @@ Y_UNIT_TEST_SUITE(TStripStringTest) { UNIT_ASSERT(s == s2); #ifndef TSTRING_IS_STD_STRING - UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + if (TStringUseCow) { + UNIT_ASSERT(s.c_str() == s2.c_str()); // Collapse() does not change the string at all + } #endif } diff --git a/util/system/context.h b/util/system/context.h index 1b690bd57c7..6faf0582c8b 100644 --- a/util/system/context.h +++ b/util/system/context.h @@ -36,7 +36,7 @@ #define USE_UCONTEXT_CONT #elif defined(_win_) #define USE_FIBER_CONT -#elif (defined(_i386_) || defined(_x86_64_) || defined(_arm64_)) && !defined(_k1om_) +#elif (defined(_x86_64_) || defined(_arm64_)) && !defined(_k1om_) #define USE_JUMP_CONT #else #define USE_UCONTEXT_CONT diff --git a/util/system/context_i686.asm b/util/system/context_i686.asm deleted file mode 100644 index 11f8cecc8e5..00000000000 --- a/util/system/context_i686.asm +++ /dev/null @@ -1,43 +0,0 @@ - [bits 32] - - %define MJB_BX 0 - %define MJB_SI 1 - %define MJB_DI 2 - %define MJB_BP 3 - %define MJB_SP 4 - %define MJB_PC 5 - %define MJB_RSP MJB_SP - %define MJB_SIZE 24 - - %define LINKAGE 4 - %define PCOFF 0 - %define PTR_SIZE 4 - - %define PARMS LINKAGE - %define JMPBUF PARMS - %define JBUF PARMS - %define VAL JBUF + PTR_SIZE - -EXPORT __mylongjmp - mov ecx, [esp + JBUF] - mov eax, [esp + VAL] - mov edx, [ecx + MJB_PC*4] - mov ebx, [ecx + MJB_BX*4] - mov esi, [ecx + MJB_SI*4] - mov edi, [ecx + MJB_DI*4] - mov ebp, [ecx + MJB_BP*4] - mov esp, [ecx + MJB_SP*4] - jmp edx - -EXPORT __mysetjmp - mov eax, [esp + JMPBUF] - mov [eax + MJB_BX*4], ebx - mov [eax + MJB_SI*4], esi - mov [eax + MJB_DI*4], edi - lea ecx, [esp + JMPBUF] - mov [eax + MJB_SP*4], ecx - mov ecx, [esp + PCOFF] - mov [eax + MJB_PC*4], ecx - mov [eax + MJB_BP*4], ebp - xor eax, eax - ret diff --git a/util/system/context_i686.h b/util/system/context_i686.h deleted file mode 100644 index 1abfd5dada5..00000000000 --- a/util/system/context_i686.h +++ /dev/null @@ -1,9 +0,0 @@ -#pragma once - -#define MJB_BP 3 -#define MJB_SP 4 -#define MJB_PC 5 -#define MJB_RBP MJB_BP -#define MJB_RSP MJB_SP - -typedef int __myjmp_buf[6]; diff --git a/util/system/context_x86_64.asm b/util/system/context_x86.S similarity index 62% rename from util/system/context_x86_64.asm rename to util/system/context_x86.S index 8bcc01e4fcf..30cf154edb3 100644 --- a/util/system/context_x86_64.asm +++ b/util/system/context_x86.S @@ -1,16 +1,31 @@ - [bits 64] +#if defined(__APPLE__) + #define CNAME(name) _ ## name +#else + #define CNAME(name) name +#endif - %define MJB_RBX 0 - %define MJB_RBP 1 - %define MJB_R12 2 - %define MJB_R13 3 - %define MJB_R14 4 - %define MJB_R15 5 - %define MJB_RSP 6 - %define MJB_PC 7 - %define MJB_SIZE (8*8) +#define EXPORT(name, id) \ + .globl CNAME(name); \ + .set CNAME(name), id ## f; \ + id: -EXPORT __mylongjmp +.text +.p2align 4 +.intel_syntax noprefix + +.code64 + +#define MJB_RBX 0 +#define MJB_RBP 1 +#define MJB_R12 2 +#define MJB_R13 3 +#define MJB_R14 4 +#define MJB_R15 5 +#define MJB_RSP 6 +#define MJB_PC 7 +#define MJB_SIZE (8 * 8) + +EXPORT(__mylongjmp, 1) mov rbx, [rdi + MJB_RBX * 8] mov rbp, [rdi + MJB_RBP * 8] mov r12, [rdi + MJB_R12 * 8] @@ -25,7 +40,7 @@ EXPORT __mylongjmp mov rsp, [rdi + MJB_RSP * 8] jmp rdx -EXPORT __mysetjmp +EXPORT(__mysetjmp, 2) mov [rdi + MJB_RBX * 8], rbx mov [rdi + MJB_RBP * 8], rbp mov [rdi + MJB_R12 * 8], r12 diff --git a/util/system/context_x86.asm b/util/system/context_x86.asm deleted file mode 100644 index e825d5d087e..00000000000 --- a/util/system/context_x86.asm +++ /dev/null @@ -1,15 +0,0 @@ -%macro EXPORT 1 - %ifdef DARWIN - global _%1 - _%1: - %else - global %1 - %1: - %endif -%endmacro - -%ifdef _x86_64_ - %include "context_x86_64.asm" -%else - %include "context_i686.asm" -%endif diff --git a/util/system/context_x86.h b/util/system/context_x86.h index 6ea066ff883..3de3799ac09 100644 --- a/util/system/context_x86.h +++ b/util/system/context_x86.h @@ -1,10 +1,6 @@ #pragma once -#if defined(_x86_64_) - #include "context_x86_64.h" -#elif defined(_i386_) - #include "context_i686.h" -#endif +#include "context_x86_64.h" #define PROGR_CNT MJB_PC #define STACK_CNT MJB_RSP diff --git a/util/system/cpu_id.h b/util/system/cpu_id.h index 32d74697298..3fa64b0f2ca 100644 --- a/util/system/cpu_id.h +++ b/util/system/cpu_id.h @@ -156,4 +156,7 @@ namespace NX86 { } // namespace NX86 +/// @return zero-terminated ASCII string stored in 'store' memory. +/// returns a meaningful result only on x86 and x86_64 architectures +/// returns an empty string on other architectures const char* CpuBrand(ui32 store[12]) noexcept; diff --git a/util/system/env.h b/util/system/env.h index 5aa4fbe400f..0f41f99ac04 100644 --- a/util/system/env.h +++ b/util/system/env.h @@ -13,7 +13,7 @@ * Search the environment list provided by the host environment for associated variable. * * @param key String identifying the name of the environmental variable to look for - * @param def String that returns if environmental variable not found by key + * @param def String that is returned if environmental variable not found by key * * @return String that is associated with the matched environment variable or the value of `def` parameter if * such variable is missing. diff --git a/util/system/fs_win.cpp b/util/system/fs_win.cpp index 65e695bee44..8d0d5b960fc 100644 --- a/util/system/fs_win.cpp +++ b/util/system/fs_win.cpp @@ -62,7 +62,7 @@ namespace NFsPrivate { WIN32_FILE_ATTRIBUTE_DATA fad; if (::GetFileAttributesExW(wname, GetFileExInfoStandard, &fad)) { if (fad.dwFileAttributes & FILE_ATTRIBUTE_READONLY) { - fad.dwFileAttributes = FILE_ATTRIBUTE_NORMAL; + fad.dwFileAttributes &= ~FILE_ATTRIBUTE_READONLY; ::SetFileAttributesW(wname, fad.dwFileAttributes); } if (fad.dwFileAttributes & FILE_ATTRIBUTE_DIRECTORY) { @@ -105,7 +105,30 @@ namespace NFsPrivate { } } } - return 0 != CreateSymbolicLinkW(lname, wname, attr != INVALID_FILE_ATTRIBUTES && (attr & FILE_ATTRIBUTE_DIRECTORY) ? SYMBOLIC_LINK_FLAG_DIRECTORY : 0); + + DWORD flags = 0; + // INVALID_FILE_ATTRIBUTES is (DWORD)-1, so every attribute bit is set in + // it, FILE_ATTRIBUTE_DIRECTORY included: a dangling link would otherwise + // come out as a directory link. + if (attr != INVALID_FILE_ATTRIBUTES && (attr & FILE_ATTRIBUTE_DIRECTORY)) { + flags |= SYMBOLIC_LINK_FLAG_DIRECTORY; + } + + // Pass SYMBOLIC_LINK_FLAG_ALLOW_UNPRIVILEGED_CREATE to allow symlink + // creation when developer mode is enabled on the current machine. + if (CreateSymbolicLinkW(lname, wname, flags | SYMBOLIC_LINK_FLAG_ALLOW_UNPRIVILEGED_CREATE)) { + return true; + } + if (::GetLastError() != ERROR_INVALID_PARAMETER) { + return false; + } + + // The flag exists since NTDDI_WIN10_RS2 (Windows 10 1703); earlier + // kernels reject the whole call with ERROR_INVALID_PARAMETER. Arcadia + // builds with _WIN32_WINNT=0x0601 (WINDOWS_VERSION_MIN in + // build/ymake_conf.py, i.e. Windows 7), so the binary has to keep + // working there: retry the way it was always done. + return 0 != CreateSymbolicLinkW(lname, wname, flags); } bool WinHardLink(const TString& existingPath, const TString& newPath) { @@ -209,6 +232,14 @@ namespace NFsPrivate { } } + static TString FromNtPath(const TString& path) { + static constexpr TStringBuf NT_PREFIX = R"(\??\)"; + if (path.StartsWith(NT_PREFIX)) { + return path.substr(NT_PREFIX.size()); + } + return path; + } + TString WinReadLink(const TString& name) { TFileHandle h = CreateFileWithUtf8Name(name, GENERIC_READ, FILE_SHARE_READ, OPEN_EXISTING, FILE_FLAG_OPEN_REPARSE_POINT | FILE_FLAG_BACKUP_SEMANTICS, true); @@ -223,11 +254,11 @@ namespace NFsPrivate { if (rdb->ReparseTag == IO_REPARSE_TAG_SYMLINK) { wchar16* str = (wchar16*)&rdb->SymbolicLinkReparseBuffer.PathBuffer[rdb->SymbolicLinkReparseBuffer.SubstituteNameOffset / sizeof(wchar16)]; size_t len = rdb->SymbolicLinkReparseBuffer.SubstituteNameLength / sizeof(wchar16); - return WideToUTF8(str, len); + return FromNtPath(WideToUTF8(str, len)); } else if (rdb->ReparseTag == IO_REPARSE_TAG_MOUNT_POINT) { wchar16* str = (wchar16*)&rdb->MountPointReparseBuffer.PathBuffer[rdb->MountPointReparseBuffer.SubstituteNameOffset / sizeof(wchar16)]; size_t len = rdb->MountPointReparseBuffer.SubstituteNameLength / sizeof(wchar16); - return WideToUTF8(str, len); + return FromNtPath(WideToUTF8(str, len)); } // this reparse point is unsupported in arcadia return TString(); diff --git a/util/system/fs_win_ut.cpp b/util/system/fs_win_ut.cpp index ce9c0d86e5f..5baea888e43 100644 --- a/util/system/fs_win_ut.cpp +++ b/util/system/fs_win_ut.cpp @@ -1,16 +1,25 @@ #include "fs.h" #include "fs_win.h" +#include + #include #include "fileapi.h" +#include "error.h" #include "file.h" #include "fstat.h" #include "win_undef.h" #include #include +#include #include +#include +#include + +#include +#include static void Touch(const TFsPath& path) { TFile file(path, CreateAlways | WrOnly); @@ -28,10 +37,159 @@ static LPCWSTR UTF8ToWCHAR(const TStringBuf str, TUtf16String& wstr) { return (const WCHAR*)wstr.data(); } +static void SetReadOnly(const TFsPath& path) { + TUtf16String wstr; + LPCWSTR wname = UTF8ToWCHAR(static_cast(path), wstr); + UNIT_ASSERT(wname); + UNIT_ASSERT(::SetFileAttributesW(wname, FILE_ATTRIBUTE_READONLY)); +} + +namespace { + // The layout FSCTL_SET_REPARSE_POINT expects for a mount point, taken from + // the same way fs_win.cpp takes it: the SDK ships no such header. + struct TMountPointReparseData { + ULONG ReparseTag; + USHORT ReparseDataLength; + USHORT Reserved; + USHORT SubstituteNameOffset; + USHORT SubstituteNameLength; + USHORT PrintNameOffset; + USHORT PrintNameLength; + wchar16 PathBuffer[1]; + }; + + constexpr size_t REPARSE_HEADER_SIZE = offsetof(TMountPointReparseData, SubstituteNameOffset); + constexpr size_t PATH_BUFFER_OFFSET = offsetof(TMountPointReparseData, PathBuffer); + static_assert(REPARSE_HEADER_SIZE == 8); + static_assert(PATH_BUFFER_OFFSET == 16); + + // A junction, unlike a symlink, asks for no privileges whatsoever, which + // makes it the one reparse point a test may rely on creating anywhere. + // Its target is always stored as an absolute NT path. + bool CreateJunction(const TString& junction, const TString& absoluteTarget) { + if (!NFsPrivate::WinMakeDirectory(junction)) { + return false; + } + + const TUtf16String substituteName = UTF8ToWide(R"(\??\)" + absoluteTarget); + const TUtf16String printName = UTF8ToWide(absoluteTarget); + const size_t substituteBytes = substituteName.size() * sizeof(wchar16); + const size_t printBytes = printName.size() * sizeof(wchar16); + + // Both names are stored NUL-terminated, one after the other. + TVector buffer( + PATH_BUFFER_OFFSET + substituteBytes + printBytes + 2 * sizeof(wchar16), + 0); + auto& data = *reinterpret_cast(buffer.data()); + data.ReparseTag = IO_REPARSE_TAG_MOUNT_POINT; + data.ReparseDataLength = static_cast(buffer.size() - REPARSE_HEADER_SIZE); + data.SubstituteNameOffset = 0; + data.SubstituteNameLength = static_cast(substituteBytes); + data.PrintNameOffset = static_cast(substituteBytes + sizeof(wchar16)); + data.PrintNameLength = static_cast(printBytes); + std::memcpy(data.PathBuffer, substituteName.data(), substituteBytes); + std::memcpy( + reinterpret_cast(data.PathBuffer) + data.PrintNameOffset, + printName.data(), + printBytes); + + TFileHandle h = NFsPrivate::CreateFileWithUtf8Name( + junction, + GENERIC_READ | GENERIC_WRITE, + FILE_SHARE_READ | FILE_SHARE_WRITE | FILE_SHARE_DELETE, + OPEN_EXISTING, + FILE_FLAG_BACKUP_SEMANTICS | FILE_FLAG_OPEN_REPARSE_POINT, + true); + if (h == INVALID_HANDLE_VALUE) { + return false; + } + + DWORD returned = 0; + return ::DeviceIoControl(h, FSCTL_SET_REPARSE_POINT, buffer.data(), + static_cast(buffer.size()), nullptr, 0, + &returned, nullptr); + } + + // Wine reports success from CreateSymbolicLinkW and from the reparse point + // ioctl without creating anything, so nothing here can be asserted there. + bool IsWine() { + if (::RegOpenKeyExA(HKEY_CURRENT_USER, R"(Software\Wine)", 0, KEY_READ, nullptr) == ERROR_SUCCESS) { + return true; + } + if (::RegOpenKeyExA(HKEY_LOCAL_MACHINE, R"(Software\Wine)", 0, KEY_READ, nullptr) == ERROR_SUCCESS) { + return true; + } + + HMODULE ntdll = ::GetModuleHandleA("ntdll.dll"); + return ntdll && ::GetProcAddress(ntdll, "wine_get_version"); + } + + bool IsProcessElevated() { + HANDLE token = nullptr; + if (!::OpenProcessToken(::GetCurrentProcess(), TOKEN_QUERY, &token)) { + return false; + } + Y_DEFER { + ::CloseHandle(token); + }; + + TOKEN_ELEVATION elevation = {}; + DWORD size = sizeof(elevation); + if (!::GetTokenInformation(token, TokenElevation, &elevation, sizeof(elevation), &size)) { + return false; + } + return elevation.TokenIsElevated != 0; + } + + // The Developer Mode switch. Returns -1 when the value is not there to read. + int DeveloperMode() { + DWORD value = 0; + DWORD size = sizeof(value); + const LONG res = ::RegGetValueA( + HKEY_LOCAL_MACHINE, + R"(SOFTWARE\Microsoft\Windows\CurrentVersion\AppModelUnlock)", + "AllowDevelopmentWithoutDevLicense", + RRF_RT_REG_DWORD, + nullptr, + &value, + &size); + return res == ERROR_SUCCESS ? static_cast(value) : -1; + } + + // Printed into the test log on purpose: it is the only way to tell from a CI + // run whether the machine it landed on lets an ordinary process create + // symlinks, and if not, which of the two reasons is to blame. + void ReportSymlinkEnvironment(bool created, int error) { + Cerr << "symlink environment: created=" << created + << " error=" << error << " (" << LastSystemErrorText(error) << ")" + << " elevated=" << IsProcessElevated() + << " developerMode=" << DeveloperMode() + << Endl; + } +} // namespace + +// Creating a symlink needs SeCreateSymbolicLinkPrivilege, which comes with +// elevation, with Developer Mode, or granted to the account outright. Report +// what the machine offers and leave the test when it offers none of it. +#define SYMLINK_OR_RETURN(target, link) \ + do { \ + const bool created = NFsPrivate::WinSymLink(target, link); \ + const int error = created ? 0 : LastSystemError(); \ + ReportSymlinkEnvironment(created, error); \ + if (!created) { \ + UNIT_ASSERT_VALUES_EQUAL(error, ERROR_PRIVILEGE_NOT_HELD); \ + return; \ + } \ + if (!NFsPrivate::WinExists(link) && IsWine()) { \ + Cerr << "wine does not support symlinks" << Endl; \ + return; \ + } \ + } while (false) + Y_UNIT_TEST_SUITE(TFsWinTest) { Y_UNIT_TEST(TestRemoveDirWithROFiles) { TFsPath dir1 = "dir1"; - NFsPrivate::WinRemove(dir1); + NFs::RemoveRecursive(dir1); UNIT_ASSERT(!NFsPrivate::WinExists(dir1)); UNIT_ASSERT(NFsPrivate::WinMakeDirectory(dir1)); @@ -39,15 +197,12 @@ Y_UNIT_TEST_SUITE(TFsWinTest) { TFsPath file1 = dir1 / "file.txt"; Touch(file1); UNIT_ASSERT(NFsPrivate::WinExists(file1)); - { - TUtf16String wstr; - LPCWSTR wname = UTF8ToWCHAR(static_cast(file1), wstr); - UNIT_ASSERT(wname); - WIN32_FILE_ATTRIBUTE_DATA fad; - fad.dwFileAttributes = FILE_ATTRIBUTE_READONLY; - ::SetFileAttributesW(wname, fad.dwFileAttributes); - } - NFsPrivate::WinRemove(dir1); + SetReadOnly(file1); + + // A read-only file may not be left in the way of a recursive removal: + // on unix the attribute of the file says nothing about deleting it, and + // windows is made to behave the same way. + NFs::RemoveRecursive(dir1); UNIT_ASSERT(!NFsPrivate::WinExists(dir1)); } @@ -58,15 +213,83 @@ Y_UNIT_TEST_SUITE(TFsWinTest) { UNIT_ASSERT(NFsPrivate::WinMakeDirectory(dir1)); UNIT_ASSERT(TFileStat(dir1).IsDir()); + SetReadOnly(dir1); + + // Dropping the attribute must not cost the directory its own type: + // what is removed here has to be removed as a directory. + UNIT_ASSERT(NFsPrivate::WinRemove(dir1)); + UNIT_ASSERT(!NFsPrivate::WinExists(dir1)); + } + + Y_UNIT_TEST(TestSymLinkToFile) { + TFsPath target = "symlink_target.txt"; + TFsPath link = "symlink.txt"; + NFsPrivate::WinRemove(link); + NFsPrivate::WinRemove(target); + Touch(target); + + SYMLINK_OR_RETURN(static_cast(target), static_cast(link)); + + UNIT_ASSERT(TFileStat(link, true).IsSymlink()); + // Following the link lands on a plain file. + UNIT_ASSERT(!TFileStat(link, false).IsSymlink()); + UNIT_ASSERT(TFileStat(link, false).IsFile()); + + // A relative target is stored as given, no spelling of our own. + UNIT_ASSERT_STRINGS_EQUAL(NFsPrivate::WinReadLink(link), target.GetPath()); { - TUtf16String wstr; - LPCWSTR wname = UTF8ToWCHAR(static_cast(dir1), wstr); - UNIT_ASSERT(wname); - WIN32_FILE_ATTRIBUTE_DATA fad; - fad.dwFileAttributes = FILE_ATTRIBUTE_READONLY; - ::SetFileAttributesW(wname, fad.dwFileAttributes); + TFile file(link, OpenExisting | RdOnly); + UNIT_ASSERT_VALUES_EQUAL(file.GetLength(), 4); } - NFsPrivate::WinRemove(dir1); - UNIT_ASSERT(!NFsPrivate::WinExists(dir1)); + + // Removing the link leaves the target alone. + UNIT_ASSERT(NFsPrivate::WinRemove(link)); + UNIT_ASSERT(NFsPrivate::WinExists(target)); + UNIT_ASSERT(NFsPrivate::WinRemove(target)); + } + + Y_UNIT_TEST(TestReadLinkOnAbsoluteSymLink) { + TFsPath target = "abs_symlink_target.txt"; + TFsPath link = "abs_symlink.txt"; + NFsPrivate::WinRemove(link); + NFsPrivate::WinRemove(target); + Touch(target); + const TString absoluteTarget = target.RealPath().GetPath(); + + SYMLINK_OR_RETURN(absoluteTarget, static_cast(link)); + + // The reparse point holds the NT spelling of the target - "\??\C:\dir\file" - + // which names nothing outside the kernel, so a win32 path is handed out. + UNIT_ASSERT_STRINGS_EQUAL(NFsPrivate::WinReadLink(link), absoluteTarget); + + UNIT_ASSERT(NFsPrivate::WinRemove(link)); + UNIT_ASSERT(NFsPrivate::WinRemove(target)); + } + + Y_UNIT_TEST(TestReadLinkOnJunction) { + TFsPath target = "junction_target"; + TFsPath junction = "junction"; + NFsPrivate::WinRemove(junction); + NFsPrivate::WinRemove(target); + UNIT_ASSERT(NFsPrivate::WinMakeDirectory(target)); + const TString absoluteTarget = target.RealPath().GetPath(); + + if (!CreateJunction(junction, absoluteTarget)) { + Cerr << "can't create junction: " + << LastSystemErrorText(LastSystemError()) << Endl; + UNIT_ASSERT(IsWine()); + return; + } + + // Needs no privileges, so this is the one check of the NT prefix that + // holds on any machine. + UNIT_ASSERT_STRINGS_EQUAL(NFsPrivate::WinReadLink(junction), absoluteTarget); + // TFileStat calls a mount point a symlink, and WinReadLink reads one. + UNIT_ASSERT(TFileStat(junction, true).IsSymlink()); + + // Removing the junction leaves the directory it points at alone. + UNIT_ASSERT(NFsPrivate::WinRemove(junction)); + UNIT_ASSERT(NFsPrivate::WinExists(target)); + UNIT_ASSERT(NFsPrivate::WinRemove(target)); } } // Y_UNIT_TEST_SUITE(TFsWinTest)