diff --git a/.github/actions/prepare_vm/action.yaml b/.github/actions/prepare_vm/action.yaml index 2e6f36d43ea..bddcf6a89ae 100644 --- a/.github/actions/prepare_vm/action.yaml +++ b/.github/actions/prepare_vm/action.yaml @@ -19,7 +19,8 @@ runs: sudo apt-get -y update sudo apt-get -y install git gdb ninja-build libidn11-dev ragel yasm libc-ares-dev libre2-dev \ rapidjson-dev zlib1g-dev libxxhash-dev libzstd-dev libsnappy-dev libgtest-dev libgmock-dev \ - libbz2-dev liblz4-dev libdouble-conversion-dev libssl-dev libstdc++-13-dev gcc-13 g++-13 + libbz2-dev liblz4-dev libdouble-conversion-dev libssl-dev libstdc++-13-dev gcc-13 g++-13 \ + unixodbc unixodbc-dev sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-13 10000 sudo update-alternatives --install /usr/bin/g++ g++ /usr/bin/g++-13 10000 @@ -28,6 +29,8 @@ runs: if: ${{ inputs.mode == 'build' }} shell: bash run: | + # Static deps are linked into the shared ODBC driver; they must be built with -fPIC. + export PIC="-DCMAKE_POSITION_INDEPENDENT_CODE=ON" # Install ccache (V=4.8.1; curl -L https://github.com/ccache/ccache/releases/download/v${V}/ccache-${V}-linux-x86_64.tar.xz | \ sudo tar -xJ -C /usr/local/bin/ --strip-components=1 --no-same-owner ccache-${V}-linux-x86_64/ccache) @@ -48,7 +51,7 @@ runs: tar -xvzf abseil-cpp-20230802.0.tar.gz cd abseil-cpp-20230802.0 mkdir build && cd build - cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_BUILD_TYPE=Release -DABSL_PROPAGATE_CXX_STD=ON .. + cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_BUILD_TYPE=Release -DABSL_PROPAGATE_CXX_STD=ON ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/absl cd ../../ @@ -59,7 +62,7 @@ runs: cd protobuf-25.0 mkdir build && cd build cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_PREFIX_PATH="${HOME}/ydb_deps/absl" -DCMAKE_BUILD_TYPE=Release \ - -Dprotobuf_BUILD_TESTS=OFF -Dprotobuf_INSTALL=ON -Dprotobuf_ABSL_PROVIDER=package .. + -Dprotobuf_BUILD_TESTS=OFF -Dprotobuf_INSTALL=ON -Dprotobuf_ABSL_PROVIDER=package ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/protobuf cd ../../ @@ -68,7 +71,7 @@ runs: wget -O grpc-1.60.2.tar.gz https://github.com/grpc/grpc/archive/refs/tags/v1.60.2.tar.gz tar -xvzf grpc-1.60.2.tar.gz && cd grpc-1.60.2 mkdir build && cd build - cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_PREFIX_PATH="${HOME}/ydb_deps/absl;${HOME}/ydb_deps/protobuf" -DCMAKE_BUILD_TYPE=Release -DCMAKE_CXX_STANDARD=17 \ + cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_PREFIX_PATH="${HOME}/ydb_deps/absl;${HOME}/ydb_deps/protobuf" -DCMAKE_BUILD_TYPE=Release -DCMAKE_CXX_STANDARD=17 ${PIC} \ -DgRPC_INSTALL=ON -DgRPC_BUILD_TESTS=OFF -DgRPC_BUILD_CSHARP_EXT=OFF \ -DgRPC_ZLIB_PROVIDER=package -DgRPC_CARES_PROVIDER=package -DgRPC_RE2_PROVIDER=package \ -DgRPC_SSL_PROVIDER=package -DgRPC_PROTOBUF_PROVIDER=package -DgRPC_ABSL_PROVIDER=package \ @@ -78,11 +81,20 @@ runs: cmake --install . --config Release --prefix ~/ydb_deps/grpc cd ../../ + # Install base64 + wget -O base64-0.5.2.tar.gz https://github.com/aklomp/base64/archive/refs/tags/v0.5.2.tar.gz + tar -xvzf base64-0.5.2.tar.gz && cd base64-0.5.2 + mkdir build && cd build + cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_BUILD_TYPE=Release ${PIC} .. + cmake --build . --config Release + cmake --install . --config Release --prefix ~/ydb_deps/base64 + cd ../../ + # Install brotli wget -O brotli-1.1.0.tar.gz https://github.com/google/brotli/archive/refs/tags/v1.1.0.tar.gz tar -xvzf brotli-1.1.0.tar.gz && cd brotli-1.1.0 mkdir build && cd build - cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_BUILD_TYPE=Release .. + cmake -G Ninja ${ENABLE_CCACHE} -DCMAKE_BUILD_TYPE=Release ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/brotli cd ../../ @@ -90,5 +102,5 @@ runs: # Clean up ccache -s sudo rm -rf llvm.sh abseil-cpp-20230802.0.tar.gz protobuf-25.0.tar.gz grpc-1.60.2.tar.gz \ - brotli-1.1.0.tar.gz abseil-cpp-20230802.0 \ - protobuf-25.0 grpc-1.60.2 brotli-1.1.0 + base64-0.5.2.tar.gz brotli-1.1.0.tar.gz abseil-cpp-20230802.0 \ + protobuf-25.0 grpc-1.60.2 base64-0.5.2 brotli-1.1.0 diff --git a/.github/scripts/run_iam_integration_tests.sh b/.github/scripts/run_iam_integration_tests.sh index 43cf89f56c8..516ee6700dd 100755 --- a/.github/scripts/run_iam_integration_tests.sh +++ b/.github/scripts/run_iam_integration_tests.sh @@ -2,7 +2,7 @@ set -euo pipefail -IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility)\.' +IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility|OdbcAuthentication)\.' IAM_CONTAINER_NAME="${IAM_CONTAINER_NAME:-ydb-iam}" IAM_CTEST_JOBS="${IAM_CTEST_JOBS:-2}" IAM_READY_ATTEMPTS="${IAM_READY_ATTEMPTS:-60}" @@ -16,7 +16,7 @@ cleanup_iam() { wait_for_iam_ydb() { for _ in $(seq 1 "${IAM_READY_ATTEMPTS}"); do if docker exec -e "YDB_TOKEN=${IAM_TOKEN}" "${IAM_CONTAINER_NAME}" /ydb \ - --endpoint grpc://localhost:2136 \ + --endpoint grpc://localhost:2236 \ --database /local \ sql -s 'select 1' >/dev/null 2>&1; then return 0 @@ -29,12 +29,22 @@ wait_for_iam_ydb() { return 1 } +provision_odbc_static_user() { + docker exec -e "YDB_TOKEN=${IAM_TOKEN}" "${IAM_CONTAINER_NAME}" /ydb \ + --endpoint grpc://localhost:2236 \ + --database /local \ + sql -s "CREATE USER odbcauth PASSWORD '12345678'" +} + trap cleanup_iam EXIT cleanup_iam docker run -d --name "${IAM_CONTAINER_NAME}" --hostname localhost \ - -p 2235:2135 -p 2236:2136 -p 28765:8765 \ + -p 2235:2235 -p 2236:2236 -p 28765:28765 \ -v /tmp/ydb_iam_certs:/ydb_certs \ + -e GRPC_TLS_PORT=2235 \ + -e GRPC_PORT=2236 \ + -e MON_PORT=28765 \ -e YDB_USE_IN_MEMORY_PDISKS=true \ -e YDB_TABLE_ENABLE_PREPARED_DDL=true \ -e YDB_ENFORCE_USER_TOKEN_REQUIREMENT=true \ @@ -42,6 +52,8 @@ docker run -d --name "${IAM_CONTAINER_NAME}" --hostname localhost \ ghcr.io/ydb-platform/local-ydb:trunk wait_for_iam_ydb +provision_odbc_static_user YDB_ENDPOINT=localhost:2236 YDB_DATABASE=/local \ +YDB_ODBC_STATIC_USER=odbcauth YDB_ODBC_STATIC_PASSWORD=12345678 \ ctest -j"${IAM_CTEST_JOBS}" --test-dir build -R "${IAM_REGEX}" --output-on-failure diff --git a/.github/workflows/coverage.yml b/.github/workflows/coverage.yml index c6d978e0df8..cd95e56c85d 100644 --- a/.github/workflows/coverage.yml +++ b/.github/workflows/coverage.yml @@ -77,7 +77,7 @@ jobs: run: | set -euo pipefail - IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility)\.' + IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility|OdbcAuthentication)\.' FLAKY_REGEX='(ManyMessages|DiscoveryHang|DescribeHang)' EXCLUDE_REGEX="${IAM_REGEX}|${FLAKY_REGEX}" diff --git a/.github/workflows/release_publish.yaml b/.github/workflows/release_publish.yaml index 9b935a766c0..6e2c31686eb 100644 --- a/.github/workflows/release_publish.yaml +++ b/.github/workflows/release_publish.yaml @@ -61,7 +61,7 @@ jobs: id: deb-package-cache-key shell: bash run: | - echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" + echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'odbc/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" - name: Restore Debian package build cache uses: actions/cache/restore@v4 diff --git a/.github/workflows/tests.yaml b/.github/workflows/tests.yaml index 54e79c108dd..b20b4167524 100644 --- a/.github/workflows/tests.yaml +++ b/.github/workflows/tests.yaml @@ -8,6 +8,7 @@ on: types: [opened, synchronize, reopened, ready_for_review] branches: - main + - odbc-driver-feature concurrency: group: ${{ github.workflow }}-${{ github.event.pull_request.number }} cancel-in-progress: true @@ -54,9 +55,15 @@ jobs: with: compiler: ${{ matrix.compiler }} - name: Test + if: github.event_name != 'pull_request' || github.base_ref == 'main' shell: bash run: | ctest -j$(nproc) --preset unit + - name: Test ODBC + if: github.event_name == 'pull_request' && github.base_ref == 'odbc-driver-feature' + shell: bash + run: | + ctest --test-dir build/odbc/tests/unit -j$(nproc) --output-on-failure - name: Package integration build shell: bash run: | @@ -132,9 +139,10 @@ jobs: tar -C build -xzf "integration-build-${{ matrix.compiler }}.tar.gz" tar -C "$HOME" -xzf "integration-deps-${{ matrix.compiler }}.tar.gz" - name: Test + if: github.event_name != 'pull_request' || github.base_ref == 'main' shell: bash run: | - IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility)\.' + IAM_REGEX='^(DriverAuth|TMetadataFixture|TJwtIamFixture|TOAuthIamFixture|OAuth_WithFacility|OdbcAuthentication)\.' YDB_VERSION=${{ matrix.ydb-version }} ctest -j2 --preset integration \ -E "${IAM_REGEX}" --output-on-failure @@ -144,8 +152,22 @@ jobs: ./.github/scripts/run_iam_integration_tests.sh ;; esac + - name: Test ODBC + if: github.event_name == 'pull_request' && github.base_ref == 'odbc-driver-feature' + shell: bash + run: | + YDB_VERSION=${{ matrix.ydb-version }} \ + ctest --test-dir build/odbc/tests/integration -j2 \ + -E '^OdbcAuthentication\.' --output-on-failure + + case '${{ matrix.ydb-version }}' in + 25.1|trunk) + ./.github/scripts/run_iam_integration_tests.sh + ;; + esac test-install: + if: github.event_name != 'pull_request' || github.base_ref == 'main' name: "Test CMake Install" concurrency: group: test-install-${{ github.ref }}-${{ matrix.compiler }} @@ -228,7 +250,7 @@ jobs: id: deb-package-cache-key shell: bash run: | - echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" + echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'odbc/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" - name: Validate dpkg-buildpackage shell: bash diff --git a/.github/workflows/warmup_cache.yaml b/.github/workflows/warmup_cache.yaml index e5864270b4d..d0d31089b06 100644 --- a/.github/workflows/warmup_cache.yaml +++ b/.github/workflows/warmup_cache.yaml @@ -94,7 +94,7 @@ jobs: id: deb-package-cache-key shell: bash run: | - echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" + echo "prefix=ubuntu-24.04-deb-packages-${{ hashFiles('CMakeLists.txt', 'cmake/**', 'contrib/**', 'include/**', 'library/**', 'odbc/**', 'plugins/**', 'scripts/build_cpack_deb_packages.sh', 'scripts/generate-debian-directory.sh', 'scripts/googleapis_deb/**', 'src/**', 'third_party/api-common-protos/**', 'tools/**', 'util/**') }}" >> "$GITHUB_OUTPUT" - name: Restore Debian package build cache id: deb-package-cache uses: actions/cache/restore@v4 diff --git a/CMakeLists.txt b/CMakeLists.txt index 72b66aad079..3ee545202cb 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -13,6 +13,11 @@ option(YDB_SDK_EXAMPLES "Build YDB C++ SDK examples" On) option(YDB_SDK_ENABLE_OTEL_METRICS "Build OpenTelemetry metrics plugin" Off) option(YDB_SDK_ENABLE_OTEL_TRACE "Build OpenTelemetry trace plugin" Off) option(YDB_CPP_SDK_SLO_USE_INSTALLED_SDK "Build only SLO workloads against an installed ydb-cpp-sdk package" Off) +option(YDB_SDK_ODBC "Build YDB ODBC driver" Off) +if (YDB_SDK_ODBC) + # ODBC driver is a shared library; static archives linked into it must be PIC. + set(CMAKE_POSITION_INDEPENDENT_CODE ON CACHE BOOL "" FORCE) +endif() set(YDB_SDK_GOOGLE_COMMON_PROTOS_TARGET "" CACHE STRING "Name of cmake target preparing google common proto library") option(YDB_SDK_USE_RAPID_JSON "Search for rapid json library in system" ON) @@ -85,6 +90,10 @@ add_subdirectory(plugins) #_ydb_sdk_validate_public_headers() +if (YDB_SDK_ODBC) + add_subdirectory(odbc) +endif() + if (YDB_SDK_EXAMPLES) add_subdirectory(examples) endif() diff --git a/CMakePresets.json b/CMakePresets.json index 2061ba6b78f..b13fdc05598 100644 --- a/CMakePresets.json +++ b/CMakePresets.json @@ -56,6 +56,7 @@ "cacheVariables": { "YDB_SDK_TESTS": "TRUE", "YDB_SDK_EXAMPLES": "TRUE", + "YDB_SDK_ODBC": "TRUE", "ARCADIA_ROOT": "..", "ARCADIA_BUILD_ROOT": "." } diff --git a/README.md b/README.md index 41d1790d193..2f1479633d3 100644 --- a/README.md +++ b/README.md @@ -57,11 +57,18 @@ match the dependency set published with that gRPC release: These pins are shared by regular CI builds, SLO workload images, and the development container. +When building with `YDB_SDK_ODBC=ON` (included in the `release-test-*` presets), the ODBC +driver is a shared library. Static dependencies installed under `~/ydb_deps` must be built +with position-independent code (`-DCMAKE_POSITION_INDEPENDENT_CODE=ON`). + ```bash sudo apt-get -y update sudo apt-get -y install git gdb ninja-build libidn11-dev ragel yasm libc-ares-dev libre2-dev \ rapidjson-dev zlib1g-dev libxxhash-dev libzstd-dev libsnappy-dev libgtest-dev libgmock-dev \ - libbz2-dev liblz4-dev libdouble-conversion-dev libssl-dev libstdc++-13-dev gcc-13 g++-13 + libbz2-dev liblz4-dev libdouble-conversion-dev libssl-dev libstdc++-13-dev gcc-13 g++-13 \ + unixodbc unixodbc-dev + +PIC="-DCMAKE_POSITION_INDEPENDENT_CODE=ON" wget https://apt.llvm.org/llvm.sh chmod u+x llvm.sh @@ -72,7 +79,7 @@ wget -O abseil-cpp-20230802.0.tar.gz https://github.com/abseil/abseil-cpp/archiv tar -xvzf abseil-cpp-20230802.0.tar.gz cd abseil-cpp-20230802.0 mkdir build && cd build -cmake -G Ninja -DCMAKE_BUILD_TYPE=Release -DABSL_PROPAGATE_CXX_STD=ON .. +cmake -G Ninja -DCMAKE_BUILD_TYPE=Release -DABSL_PROPAGATE_CXX_STD=ON ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/absl cd ../../ @@ -83,7 +90,7 @@ tar -xvzf protobuf-25.0.tar.gz cd protobuf-25.0 mkdir build && cd build cmake -G Ninja -DCMAKE_PREFIX_PATH="$HOME/ydb_deps/absl" -DCMAKE_BUILD_TYPE=Release \ - -Dprotobuf_BUILD_TESTS=OFF -Dprotobuf_INSTALL=ON -Dprotobuf_ABSL_PROVIDER=package .. + -Dprotobuf_BUILD_TESTS=OFF -Dprotobuf_INSTALL=ON -Dprotobuf_ABSL_PROVIDER=package ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/protobuf cd ../../ @@ -92,7 +99,7 @@ cd ../../ wget -O grpc-1.60.2.tar.gz https://github.com/grpc/grpc/archive/refs/tags/v1.60.2.tar.gz tar -xvzf grpc-1.60.2.tar.gz && cd grpc-1.60.2 mkdir build && cd build -cmake -G Ninja -DCMAKE_PREFIX_PATH="$HOME/ydb_deps/absl;$HOME/ydb_deps/protobuf" -DCMAKE_BUILD_TYPE=Release -DCMAKE_CXX_STANDARD=17 \ +cmake -G Ninja -DCMAKE_PREFIX_PATH="${HOME}/ydb_deps/absl;${HOME}/ydb_deps/protobuf" -DCMAKE_BUILD_TYPE=Release -DCMAKE_CXX_STANDARD=17 ${PIC} \ -DgRPC_INSTALL=ON -DgRPC_BUILD_TESTS=OFF -DgRPC_BUILD_CSHARP_EXT=OFF \ -DgRPC_ZLIB_PROVIDER=package -DgRPC_CARES_PROVIDER=package -DgRPC_RE2_PROVIDER=package \ -DgRPC_SSL_PROVIDER=package -DgRPC_PROTOBUF_PROVIDER=package -DgRPC_ABSL_PROVIDER=package \ @@ -106,7 +113,7 @@ cd ../../ wget -O base64-0.5.2.tar.gz https://github.com/aklomp/base64/archive/refs/tags/v0.5.2.tar.gz tar -xvzf base64-0.5.2.tar.gz && cd base64-0.5.2 mkdir build && cd build -cmake -G Ninja -DCMAKE_BUILD_TYPE=Release .. +cmake -G Ninja -DCMAKE_BUILD_TYPE=Release ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/base64 cd ../../ @@ -115,7 +122,7 @@ cd ../../ wget -O brotli-1.1.0.tar.gz https://github.com/google/brotli/archive/refs/tags/v1.1.0.tar.gz tar -xvzf brotli-1.1.0.tar.gz && cd brotli-1.1.0 mkdir build && cd build -cmake -G Ninja -DCMAKE_BUILD_TYPE=Release \ +cmake -G Ninja -DCMAKE_BUILD_TYPE=Release ${PIC} \ -DCMAKE_INSTALL_PREFIX="$HOME/ydb_deps/brotli" .. cmake --build . --config Release cmake --install . --config Release @@ -125,7 +132,7 @@ cd ../../ wget -O jwt-cpp-0.7.0.tar.gz https://github.com/Thalhammer/jwt-cpp/archive/refs/tags/v0.7.0.tar.gz tar -xvzf jwt-cpp-0.7.0.tar.gz && cd jwt-cpp-0.7.0 mkdir build && cd build -cmake -G Ninja -DCMAKE_BUILD_TYPE=Release .. +cmake -G Ninja -DCMAKE_BUILD_TYPE=Release ${PIC} .. cmake --build . --config Release cmake --install . --config Release --prefix ~/ydb_deps/jwt-cpp cd ../../ @@ -241,12 +248,17 @@ wget "${BASE}/libydb-cpp-dev_${TAG#v}_amd64.deb" wget "${BASE}/libydb-cpp-iam-dev_${TAG#v}_amd64.deb" wget "${BASE}/libydb-cpp-otel-metrics-dev_${TAG#v}_amd64.deb" wget "${BASE}/libydb-cpp-otel-tracing-dev_${TAG#v}_amd64.deb" +# Optional ODBC driver: +wget "${BASE}/ydb-odbc_${TAG#v}_amd64.deb" sudo apt-get update sudo apt-get install -y \ ./yandex-googleapis-api-common-protos-*.deb \ ./libydb-cpp-dev_*.deb ./libydb-cpp-iam-dev_*.deb \ - ./libydb-cpp-otel-metrics-dev_*.deb ./libydb-cpp-otel-tracing-dev_*.deb + ./libydb-cpp-otel-metrics-dev_*.deb ./libydb-cpp-otel-tracing-dev_*.deb \ + ./ydb-odbc_*.deb + +odbcinst -q -d -n YDB ``` After installation, use the SDK in your CMake project: diff --git a/cmake/PackSDK.cmake b/cmake/PackSDK.cmake index c204e7a49e3..c93a48c988a 100644 --- a/cmake/PackSDK.cmake +++ b/cmake/PackSDK.cmake @@ -18,6 +18,9 @@ set(CPACK_RESOURCE_FILE_LICENSE "${YDB_SDK_SOURCE_DIR}/LICENSE") set(CPACK_DEB_COMPONENT_INSTALL ON) set(CPACK_COMPONENTS_ALL libydb-cpp libydb-cpp-iam libydb-cpp-otel-metrics libydb-cpp-otel-tracing) +if (YDB_SDK_ODBC) + list(APPEND CPACK_COMPONENTS_ALL ydb-odbc) +endif() set(CPACK_DEBIAN_LIBYDB_CPP_PACKAGE_NAME "libydb-cpp-dev") set(CPACK_DEBIAN_LIBYDB_CPP_PACKAGE_DEPENDS @@ -34,6 +37,18 @@ set(CPACK_DEBIAN_LIBYDB_CPP_OTEL_TRACING_PACKAGE_NAME "libydb-cpp-otel-tracing-d set(CPACK_DEBIAN_LIBYDB_CPP_OTEL_TRACING_PACKAGE_DEPENDS "libydb-cpp-dev (= ${YDB_SDK_VERSION}), libydb-cpp-otel-metrics-dev (= ${YDB_SDK_VERSION})") +if (YDB_SDK_ODBC) + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_NAME" "ydb-odbc") + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_DEPENDS" "odbcinst") + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_SHLIBDEPS" ON) + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_CONTROL_EXTRA" + "${YDB_ODBC_DEBIAN_CONTROL_EXTRA}") + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_CONTROL_STRICT_PERMISSION" ON) + set("CPACK_DEBIAN_YDB-ODBC_PACKAGE_SECTION" "database") + set("CPACK_DEBIAN_YDB-ODBC_DESCRIPTION" + "YDB ODBC driver\n Shared ODBC driver and unixODBC registration for YDB.") +endif() + foreach(component IN ITEMS libydb-cpp libydb-cpp-iam libydb-cpp-otel-metrics libydb-cpp-otel-tracing) string(TOUPPER "${component}" component_upper) string(REPLACE "-" "_" component_var "${component_upper}") diff --git a/cmake/common.cmake b/cmake/common.cmake index b40eee34324..083c8a2c3d4 100644 --- a/cmake/common.cmake +++ b/cmake/common.cmake @@ -110,7 +110,7 @@ function(generate_enum_serilization Tgt Input) endfunction() function(add_global_library_for TgtName MainName) - add_library(${TgtName} STATIC ${ARGN}) + _ydb_sdk_add_library(${TgtName} STATIC ${ARGN}) if(APPLE) target_link_options(${MainName} INTERFACE "SHELL:-Wl,-force_load,$${TgtName}>") else() @@ -196,7 +196,7 @@ endfunction() function(_ydb_sdk_add_library Tgt) cmake_parse_arguments(ARG - "INTERFACE" "" "" + "INTERFACE;OBJECT;SHARED" "" "" ${ARGN} ) @@ -206,7 +206,12 @@ function(_ydb_sdk_add_library Tgt) set(libraryMode "INTERFACE") set(includeMode "INTERFACE") endif() - + if (ARG_OBJECT) + set(libraryMode "OBJECT") + endif() + if (ARG_SHARED) + set(libraryMode "SHARED") + endif() add_library(${Tgt} ${libraryMode}) target_include_directories(${Tgt} ${includeMode} $ @@ -217,6 +222,7 @@ function(_ydb_sdk_add_library Tgt) YDB_SDK_OSS ) _ydb_sdk_apply_coverage(${Tgt}) + set_property(TARGET ${Tgt} PROPERTY POSITION_INDEPENDENT_CODE ON) endfunction() diff --git a/cmake/external_libs.cmake b/cmake/external_libs.cmake index 469df6000ca..e88fe817577 100644 --- a/cmake/external_libs.cmake +++ b/cmake/external_libs.cmake @@ -197,6 +197,10 @@ if (YDB_SDK_ENABLE_OTEL_METRICS OR YDB_SDK_ENABLE_OTEL_TRACE) set(CMAKE_INSTALL_DEFAULT_COMPONENT_NAME "${_ydb_sdk_saved_install_component}") endif() +if (YDB_SDK_ODBC) + find_package(ODBC REQUIRED) +endif() + # RapidJSON if (YDB_SDK_USE_RAPID_JSON) find_package(RapidJSON REQUIRED) diff --git a/cmake/testing.cmake b/cmake/testing.cmake index 999e6a596d6..a45477df1a7 100644 --- a/cmake/testing.cmake +++ b/cmake/testing.cmake @@ -79,20 +79,32 @@ function(add_ydb_test) if (YDB_TEST_GTEST) set(env_vars "") - foreach(env_var ${YDB_TEST_ENV}) + foreach(env_var IN LISTS YDB_TEST_ENV) list(APPEND env_vars "ENVIRONMENT") - list(APPEND env_vars ${env_var}) + list(APPEND env_vars "${env_var}") endforeach() - gtest_discover_tests(${YDB_TEST_NAME} EXTRA_ARGS ${YDB_TEST_TEST_ARG} WORKING_DIRECTORY ${YDB_TEST_WORKING_DIRECTORY} PROPERTIES - LABELS ${YDB_TEST_LABELS} ENVIRONMENT "YDB_TEST_ROOT=sdk_tests" ${env_vars} ) + # Discovered tests only exist when CTest loads this directory. Assign + # labels from a second include so a semicolon-separated label list stays a + # single property value rather than becoming extra property/value pairs. + if (YDB_TEST_LABELS) + set(test_labels_file + "${CMAKE_CURRENT_BINARY_DIR}/${YDB_TEST_NAME}_labels.cmake") + string(CONCAT test_labels_content + "if(DEFINED ${YDB_TEST_NAME}_TESTS)\n" + " set_tests_properties(\${${YDB_TEST_NAME}_TESTS} PROPERTIES LABELS \"${YDB_TEST_LABELS}\")\n" + "endif()\n") + file(GENERATE OUTPUT "${test_labels_file}" CONTENT "${test_labels_content}") + set_property(DIRECTORY APPEND PROPERTY TEST_INCLUDE_FILES "${test_labels_file}") + endif() + target_link_libraries(${YDB_TEST_NAME} PRIVATE GTest::gtest_main GTest::gmock_main @@ -113,7 +125,7 @@ function(add_ydb_test) cpp-testing-unittest_main ) - set_tests_properties(${YDB_TEST_NAME} PROPERTIES LABELS ${YDB_TEST_LABELS}) + set_tests_properties(${YDB_TEST_NAME} PROPERTIES LABELS "${YDB_TEST_LABELS}") set_tests_properties(${YDB_TEST_NAME} PROPERTIES ENVIRONMENT "YDB_TEST_ROOT=sdk_tests") if (YDB_TEST_ENV) set_tests_properties(${YDB_TEST_NAME} PROPERTIES ENVIRONMENT ${YDB_TEST_ENV}) @@ -122,3 +134,37 @@ function(add_ydb_test) vcs_info(${YDB_TEST_NAME}) endfunction() + +if (YDB_SDK_ODBC) + function(add_odbc_test) + set(opts "") + set(oneval_args NAME WORKING_DIRECTORY OUTPUT_DIRECTORY) + set(multival_args SOURCES LINK_LIBRARIES LABELS) + cmake_parse_arguments(ODBC_TEST + "${opts}" + "${oneval_args}" + "${multival_args}" + ${ARGN} + ) + + add_ydb_test(GTEST + NAME ${ODBC_TEST_NAME} + SOURCES ${ODBC_TEST_SOURCES} + LINK_LIBRARIES + ${ODBC_TEST_LINK_LIBRARIES} + ODBC::ODBC + LABELS + integration + ${ODBC_TEST_LABELS} + ) + + target_compile_definitions(${ODBC_TEST_NAME} + PRIVATE + ODBC_DRIVER_PATH="$" + ODBC_TEST_ODBCINI="${CMAKE_BINARY_DIR}/odbc/odbc.ini" + ODBC_TEST_ODBCSYSINI="${CMAKE_BINARY_DIR}/odbc" + ) + + add_dependencies(${ODBC_TEST_NAME} ydb-odbc) + endfunction() +endif() diff --git a/odbc/CMakeLists.txt b/odbc/CMakeLists.txt new file mode 100644 index 00000000000..f5416b7b606 --- /dev/null +++ b/odbc/CMakeLists.txt @@ -0,0 +1,121 @@ +add_library(ydb-odbc SHARED + src/utils/attr.cpp + src/utils/escape.cpp + src/utils/sql_type_map.cpp + src/utils/param_rewrite.cpp + src/utils/type_info_rows.cpp + src/utils/cursor.cpp + src/utils/types.cpp + src/utils/util.cpp + src/utils/status_util.cpp + src/utils/convert.cpp + src/utils/error_manager.cpp + src/odbc_driver.cpp + src/connection_attr.cpp + src/connection_config.cpp + src/connection.cpp + src/statement_attr.cpp + src/statement.cpp + src/statement_metadata.cpp + src/environment.cpp + src/metadata.cpp + src/descriptor.cpp +) + +target_include_directories(ydb-odbc + PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR}/include + ${CMAKE_CURRENT_SOURCE_DIR}/src + ${ODBC_INCLUDE_DIRS} +) + +target_link_libraries(ydb-odbc + PRIVATE + YDB-CPP-SDK::Query + YDB-CPP-SDK::Table + YDB-CPP-SDK::Scheme + YDB-CPP-SDK::Driver + YDB-CPP-SDK::Credentials + YDB-CPP-SDK::Helpers + YDB-CPP-SDK::Iam + ODBC::ODBC + odbcinst +) + +set_target_properties(ydb-odbc PROPERTIES + POSITION_INDEPENDENT_CODE ON +) + +include(GNUInstallDirs) + +set(YDB_ODBC_INSTALL_LIBDIR "${CMAKE_INSTALL_LIBDIR}" CACHE STRING + "Directory where the YDB ODBC shared library is installed") +set(YDB_ODBC_INSTALL_DATADIR "${CMAKE_INSTALL_DATAROOTDIR}/ydb-odbc" CACHE STRING + "Directory where the YDB ODBC driver registration template is installed") + +if (IS_ABSOLUTE "${YDB_ODBC_INSTALL_LIBDIR}") + set(YDB_ODBC_DRIVER_INSTALL_DIR "${YDB_ODBC_INSTALL_LIBDIR}") +else() + set(YDB_ODBC_DRIVER_INSTALL_DIR + "${CMAKE_INSTALL_PREFIX}/${YDB_ODBC_INSTALL_LIBDIR}") +endif() + +if (IS_ABSOLUTE "${YDB_ODBC_INSTALL_DATADIR}") + set(YDB_ODBC_DRIVER_TEMPLATE_DIR "${YDB_ODBC_INSTALL_DATADIR}") +else() + set(YDB_ODBC_DRIVER_TEMPLATE_DIR + "${CMAKE_INSTALL_PREFIX}/${YDB_ODBC_INSTALL_DATADIR}") +endif() + +file(GENERATE + OUTPUT "${CMAKE_CURRENT_BINARY_DIR}/odbcinst.ini" + CONTENT "[YDB] +Description=YDB ODBC Driver +Driver=$ +Setup=$ +" +) + +set(YDB_ODBC_DRIVER_PATH + "${YDB_ODBC_DRIVER_INSTALL_DIR}/libydb-odbc${CMAKE_SHARED_LIBRARY_SUFFIX}") +configure_file( + "${CMAKE_CURRENT_SOURCE_DIR}/odbcinst.ini.in" + "${CMAKE_CURRENT_BINARY_DIR}/ydb-odbc-odbcinst.ini" + @ONLY +) + +set(YDB_ODBC_DRIVER_TEMPLATE_PATH + "${YDB_ODBC_DRIVER_TEMPLATE_DIR}/odbcinst.ini") +file(MAKE_DIRECTORY "${CMAKE_CURRENT_BINARY_DIR}/debian") +configure_file( + "${CMAKE_CURRENT_SOURCE_DIR}/packaging/postinst.in" + "${CMAKE_CURRENT_BINARY_DIR}/debian/postinst" + @ONLY +) +configure_file( + "${CMAKE_CURRENT_SOURCE_DIR}/packaging/prerm.in" + "${CMAKE_CURRENT_BINARY_DIR}/debian/prerm" + @ONLY +) +set(YDB_ODBC_DEBIAN_CONTROL_EXTRA + "${CMAKE_CURRENT_BINARY_DIR}/debian/postinst;${CMAKE_CURRENT_BINARY_DIR}/debian/prerm" + CACHE INTERNAL "Debian control scripts for the ydb-odbc package") + +install(FILES "${CMAKE_CURRENT_BINARY_DIR}/ydb-odbc-odbcinst.ini" + DESTINATION "${YDB_ODBC_INSTALL_DATADIR}" + RENAME odbcinst.ini + COMPONENT ydb-odbc +) + +install(TARGETS ydb-odbc + LIBRARY DESTINATION "${YDB_ODBC_INSTALL_LIBDIR}" + COMPONENT ydb-odbc +) + +if (YDB_SDK_EXAMPLES) + add_subdirectory(examples) +endif() + +if (YDB_SDK_TESTS) + add_subdirectory(tests) +endif() diff --git a/odbc/README.md b/odbc/README.md new file mode 100644 index 00000000000..b4a9b634837 --- /dev/null +++ b/odbc/README.md @@ -0,0 +1,163 @@ +# YDB ODBC Driver + +ODBC driver for YDB. + +## Requirements + +- CMake 3.10 or higher +- C/C++ compiler with C11 and C++20 support +- YDB C++ SDK (build with `YDB_SDK_ODBC=ON`) +- unixODBC development packages (`unixodbc`, `unixodbc-dev` on Debian/Ubuntu) + +Static dependencies under `~/ydb_deps` must be built with +`-DCMAKE_POSITION_INDEPENDENT_CODE=ON` when linking the shared ODBC driver. See the +main [README](../README.md) dependency install section. + +## Build + +```bash +cmake --preset release-test-clang +cmake --build build --target ydb-odbc -j$(nproc) +``` + +The shared library is produced as `build/odbc/libydb-odbc.so`. + +## Install + +```bash +cmake --install build --prefix /usr/local +sudo odbcinst -i -d -f /usr/local/share/ydb-odbc/odbcinst.ini +``` + +This installs `libydb-odbc` and its unixODBC registration template. The +`ydb-odbc` Debian package runs `odbcinst` automatically during installation +and unregisters the driver when the package is removed. `odbc.ini` is not +installed or modified — create your own DSN (see below). + +## Configuration + +For `SQLConnect("YDB", ...)`, `isql -v YDB`, or `Driver=YDB`. + +**`odbcinst.ini`** — driver registration template (generated on build/install). +Section `[YDB]` is the driver name used as `Driver=YDB` in connection strings +and DSNs. `Driver` and `Setup` are the full path to `libydb-odbc.so`. Register +the template with `odbcinst -i -d -f`; the Debian package does this for you. + +```ini +[YDB] +Description=YDB ODBC Driver +Driver=/path/to/libydb-odbc.so +Setup=/path/to/libydb-odbc.so +``` + +**`odbc.ini`** — DSN named `YDB`. In section `[YDB]`: `Driver` is the registered driver name, `Server` is the YDB endpoint, `Database` is the database path. Use `/etc/odbc.ini` or set `ODBCINI` to your file path. + +```ini +[ODBC Data Sources] +YDB=YDB ODBC Driver + +[YDB] +Driver=YDB +Server=localhost:2136 +Database=/local +AuthMode=Anonymous +``` + +`SQLDriverConnect` may also combine a DSN with explicit attributes. Values in +the connection string take precedence over values from the DSN. The user name +and password passed to `SQLConnect` take precedence over `User` and `Password` +in the DSN. + +### Connection attributes + +| Attribute | Meaning | +| --- | --- | +| `Endpoint` | YDB endpoint. `Server` is an alias. A `grpc://` prefix forces a plaintext connection; `grpcs://` enables TLS. | +| `Database` | YDB database path. | +| `DSN` | DSN section to load before applying the remaining connection-string attributes. | +| `AuthMode` | `Anonymous`, `Token`, `Static`, `Metadata`, `ServiceAccount`, `OAuth2`, or `Environment`. Values are case-insensitive. | +| `Token` | Access token for `Token` mode. `AccessToken` is an alias. | +| `User`, `Password` | Credentials for `Static` mode. `UID` and `PWD` are aliases. | +| `MetadataHost`, `MetadataPort` | Optional metadata service address for `Metadata` mode. | +| `ServiceAccountKeyFile` | Path to a service-account JSON key for `ServiceAccount` mode. `SaFile` is an alias. | +| `OAuth2KeyFile` | Path to an OAuth 2.0 token-exchange configuration file for `OAuth2` mode. | +| `IamEndpoint` | IAM gRPC endpoint for service-account authentication, or HTTP token endpoint override for OAuth 2.0 token exchange. | +| `RootCertificate` | Path to a PEM root-certificate file. `CaFile` is an alias. | +| `ClientCertificate`, `ClientPrivateKey` | Paths to the PEM client certificate and private key. They must be specified together. | + +If `AuthMode` is omitted, the driver infers it from exactly one credential +family (`Token`, static user/password, metadata settings, service-account key, +or OAuth 2.0 key). With no credential attributes it uses `Anonymous`. Conflicting +families and incomplete credentials are rejected with SQLSTATE `28000`. +`Environment` uses the SDK's standard `YDB_*_CREDENTIALS` variables. + +Unrecognized connection-string attributes are ignored after reporting SQLSTATE +`01S00`; `SQLDriverConnect` completes with `SQL_SUCCESS_WITH_INFO`. This allows +ODBC applications to supply tool-specific attributes such as `APP` or `WSID`. + +Certificate attributes contain file paths, not inline PEM. The driver reads the +files while establishing the ODBC connection. Supplying certificates enables +TLS; certificates cannot be combined with an explicitly plaintext `grpc://` +endpoint. + +Examples: + +```text +Driver=YDB;Endpoint=grpcs://ydb.example.net:2135;Database=/production;AuthMode=Token;Token=... +DSN=YDB;AuthMode=Static;UID=app;PWD=secret +Driver=YDB;Endpoint=localhost:2136;Database=/local;AuthMode=ServiceAccount;SaFile=/run/secrets/sa.json;IamEndpoint=grpc://localhost:4284 +``` + +## Usage + +Example of connecting via isql: +```bash +isql -v YDB +``` + +Example usage in C: +```c +SQLHENV env; +SQLHDBC dbc; +SQLHSTMT stmt; + +// Initialize environment +SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env); +SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0); + +// Connect +SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc); +SQLConnect(dbc, (SQLCHAR*)"YDB", SQL_NTS, + (SQLCHAR*)"", SQL_NTS, + (SQLCHAR*)"", SQL_NTS); + +// Execute query +SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt); +SQLExecDirect(stmt, (SQLCHAR*)"SELECT * FROM mytable", SQL_NTS); + +// Cleanup +SQLFreeHandle(SQL_HANDLE_STMT, stmt); +SQLDisconnect(dbc); +SQLFreeHandle(SQL_HANDLE_DBC, dbc); +SQLFreeHandle(SQL_HANDLE_ENV, env); +``` + +Alternatively, use `SQLDriverConnect` with a connection string (does not require DSN in odbc.ini): +```c +SQLCHAR connStr[] = "Driver=YDB;Endpoint=localhost:2136;Database=/local"; +SQLDriverConnect(dbc, NULL, connStr, SQL_NTS, NULL, 0, NULL, SQL_DRIVER_NOPROMPT); +``` + +For `INSERT`, `UPDATE`, `DELETE`, `UPSERT`, and `REPLACE`, `SQLRowCount` +returns the affected-row count reported by YDB query statistics. Counts from +executed parameter-array entries are summed; ignored entries are not counted. +For statements without an applicable count, it returns `-1`. + +## Parameters + +`?` placeholders are rewritten to `$p1`, `$p2`, ... with auto-generated `DECLARE $pN AS ?;` +from `SQLBindParameter` types. YDB-native `$pN` syntax also works. + +## License + +Apache License 2.0 diff --git a/odbc/examples/CMakeLists.txt b/odbc/examples/CMakeLists.txt new file mode 100644 index 00000000000..88b1f27cc60 --- /dev/null +++ b/odbc/examples/CMakeLists.txt @@ -0,0 +1,2 @@ +add_subdirectory(basic) +add_subdirectory(scheme) diff --git a/odbc/examples/basic/CMakeLists.txt b/odbc/examples/basic/CMakeLists.txt new file mode 100644 index 00000000000..b99d1175f43 --- /dev/null +++ b/odbc/examples/basic/CMakeLists.txt @@ -0,0 +1,14 @@ +add_executable(odbc_basic + main.cpp +) + +target_link_libraries(odbc_basic + PRIVATE + ODBC::ODBC +) +target_compile_definitions(odbc_basic + PRIVATE + ODBC_DRIVER_PATH="$" +) + +add_dependencies(odbc_basic ydb-odbc) diff --git a/odbc/examples/basic/main.cpp b/odbc/examples/basic/main.cpp new file mode 100644 index 00000000000..8084e32f3d1 --- /dev/null +++ b/odbc/examples/basic/main.cpp @@ -0,0 +1,132 @@ +#include +#include + +#include + +void PrintOdbcError(SQLSMALLINT handleType, SQLHANDLE handle) { + SQLCHAR sqlState[6] = {0}; + SQLINTEGER nativeError = 0; + SQLCHAR message[256] = {0}; + SQLSMALLINT textLength = 0; + SQLGetDiagRec(handleType, handle, 1, sqlState, &nativeError, message, sizeof(message), &textLength); + std::cerr << "ODBC error: [" << sqlState << "] " << message << std::endl; +} + +int main() { + SQLHENV henv = nullptr; + SQLHDBC hdbc = nullptr; + SQLHSTMT hstmt = nullptr; + SQLRETURN ret; + + std::cout << "1. Allocating environment handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &henv); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating environment handle" << std::endl; + return 1; + } + SQLSetEnvAttr(henv, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0); + + std::cout << "2. Allocating connection handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_DBC, henv, &hdbc); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating connection handle" << std::endl; + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "3. Building connection string" << std::endl; + std::string connStr = "Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;"; + SQLCHAR outConnStr[1024] = {0}; + SQLSMALLINT outConnStrLen = 0; + + std::cout << "4. Connecting with SQLDriverConnect" << std::endl; + ret = SQLDriverConnect(hdbc, NULL, (SQLCHAR*)connStr.c_str(), SQL_NTS, + outConnStr, sizeof(outConnStr), &outConnStrLen, SQL_DRIVER_COMPLETE); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error connecting with SQLDriverConnect" << std::endl; + PrintOdbcError(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "5. Allocating statement handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_STMT, hdbc, &hstmt); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating statement handle" << std::endl; + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "6. Executing query" << std::endl; + SQLCHAR query[] = R"( + DECLARE $p1 AS Int64?; + SELECT id, data from test_table WHERE id == $p1; + )"; + + int64_t paramValue = 1; + SQLLEN paramInd = 0; + ret = SQLBindParameter(hstmt, 1, SQL_PARAM_INPUT, SQL_C_SBIGINT, SQL_BIGINT, 0, 0, ¶mValue, 0, ¶mInd); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error binding parameter" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + SQLFreeHandle(SQL_HANDLE_STMT, hstmt); + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + ret = SQLExecDirect(hstmt, query, SQL_NTS); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error executing query" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + SQLFreeHandle(SQL_HANDLE_STMT, hstmt); + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "7. Fetching result" << std::endl; + + SQLLEN ind = 0; + int value1 = 0; + if (SQLBindCol(hstmt, 1, SQL_C_SLONG, &value1, 0, &ind) != SQL_SUCCESS) { + std::cerr << "Error binding column 1" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + SQLCHAR value2[1024] = {0}; + if (SQLBindCol(hstmt, 2, SQL_C_CHAR, &value2, 1024, &ind) != SQL_SUCCESS) { + std::cerr << "Error binding column 2" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + while ((ret = SQLFetch(hstmt)) == SQL_SUCCESS || ret == SQL_SUCCESS_WITH_INFO) { + if (ret != SQL_SUCCESS) { + std::cerr << "Error fetching result" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + std::cout << "Result column 1: " << value1 << std::endl; + std::cout << "Result column 2: " << value2 << std::endl; + + std::cout << "--------------------------------" << std::endl; + } + + std::cout << "8. Cleaning up" << std::endl; + + SQLCloseCursor(hstmt); + SQLFreeHandle(SQL_HANDLE_STMT, hstmt); + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + + return 0; +} diff --git a/odbc/examples/erlang_client/Makefile b/odbc/examples/erlang_client/Makefile new file mode 100644 index 00000000000..24bafca0f8a --- /dev/null +++ b/odbc/examples/erlang_client/Makefile @@ -0,0 +1,50 @@ +.PHONY: all compile run run-shell clean distclean check help + +CONN ?= Driver=YDB;Endpoint=localhost:2136;Database=/local; +ERLC ?= erlc +ERL ?= erl + +all: compile + +help: + @echo "YDB Series Example - Erlang ODBC" + @echo "" + @echo "Targets:" + @echo " make compile Compile Erlang modules" + @echo " make run Run the example" + @echo " make run CONN='...' Run with a custom connection string" + @echo " make run-shell Start Erlang shell with compiled modules" + @echo " make check Check Erlang ODBC availability" + @echo " make clean Remove compiled files" + @echo "" + @echo "Examples:" + @echo " make run" + @echo " make run CONN=\"Driver=YDB;Endpoint=myhost:2136;Database=/mydb;\"" + +prepare: + @mkdir -p ebin + +compile: prepare + @echo "Compiling Erlang modules..." + $(ERLC) -o ebin src/*.erl + @echo "Done." + +run: compile + @echo "Running YDB Series Example..." + $(ERL) -pa ebin -noshell -eval 'case ydb_series_client:run("$(CONN)") of ok -> halt(0); _ -> halt(1) end.' + +run-shell: compile + @echo "Starting Erlang shell with ydb_series_client..." + $(ERL) -pa ebin + +check: + @echo "Checking ODBC support in Erlang..." + @$(ERL) -noshell -eval 'application:load(odbc), io:format("~p~n", [application:start(odbc)]), halt().' + +clean: + @rm -rf ebin/*.beam + @rm -rf *.beam + @rm -rf erl_crash.dump + +distclean: clean + @rm -rf ebin diff --git a/odbc/examples/erlang_client/README.md b/odbc/examples/erlang_client/README.md new file mode 100644 index 00000000000..3a782ba9444 --- /dev/null +++ b/odbc/examples/erlang_client/README.md @@ -0,0 +1,30 @@ +# Erlang ODBC Series Example + +Minimal Erlang client for YDB ODBC. Mirrors the main scenario of the C++ `basic_example`: creates `series`, `seasons`, and `episodes` tables, fills them with test data, runs several queries, and drops the tables. + +## Requirements + +- Erlang/OTP with the `odbc` module +- unixODBC +- Built and registered YDB ODBC driver +- Running YDB instance reachable via the connection string + +Verify driver registration: + +```bash +odbcinst -q -d +``` + +## Running + +By default, `Driver=YDB;Endpoint=localhost:2136;Database=/local;` is used. + +```bash +make run +``` + +With a different connection string: + +```bash +make run CONN='...' +``` diff --git a/odbc/examples/erlang_client/src/sample_data.erl b/odbc/examples/erlang_client/src/sample_data.erl new file mode 100644 index 00000000000..f4f0ef654ce --- /dev/null +++ b/odbc/examples/erlang_client/src/sample_data.erl @@ -0,0 +1,59 @@ +-module(sample_data). +-export([series/0, seasons/0, episodes/0]). + +series() -> + [ + [1, "IT Crowd", "The IT Crowd is a British sitcom by Channel 4.", days_from_date({2006, 2, 3})], + [2, "Silicon Valley", "Silicon Valley is an American comedy series.", days_from_date({2014, 4, 6})] + ]. + +seasons() -> + [ + [1, 1, "Season 1", days_from_date({2006, 2, 3}), days_from_date({2006, 5, 5})], + [1, 2, "Season 2", days_from_date({2007, 8, 24}), days_from_date({2007, 11, 16})], + [1, 3, "Season 3", days_from_date({2008, 11, 21}), days_from_date({2008, 12, 26})], + [1, 4, "Season 4", days_from_date({2010, 6, 25}), days_from_date({2010, 7, 30})], + [2, 1, "Season 1", days_from_date({2014, 4, 6}), days_from_date({2014, 6, 15})], + [2, 2, "Season 2", days_from_date({2015, 4, 12}), days_from_date({2015, 6, 14})], + [2, 3, "Season 3", days_from_date({2016, 4, 24}), days_from_date({2016, 6, 26})], + [2, 4, "Season 4", days_from_date({2017, 4, 23}), days_from_date({2017, 6, 25})], + [2, 5, "Season 5", days_from_date({2018, 3, 25}), days_from_date({2018, 5, 13})], + [2, 6, "Season 6", days_from_date({2019, 10, 27}), days_from_date({2019, 12, 8})] + ]. + +episodes() -> + [ + [1, 1, 1, "Yesterday's Jam", days_from_date({2006, 2, 3})], + [1, 1, 2, "Calamity Jen", days_from_date({2006, 2, 10})], + [1, 1, 3, "Fifty-Fifty", days_from_date({2006, 2, 17})], + [1, 1, 4, "The Red Door", days_from_date({2006, 2, 24})], + [1, 1, 5, "The Haunting of Bill Crouse", days_from_date({2006, 3, 3})], + [1, 1, 6, "Aunt Irma Visits", days_from_date({2006, 3, 10})], + [1, 2, 1, "The Work Outing", days_from_date({2007, 8, 24})], + [1, 2, 2, "Return of the Golden Child", days_from_date({2007, 8, 31})], + [1, 2, 3, "Moss and the German", days_from_date({2007, 9, 7})], + [2, 1, 1, "Minimum Viable Product", days_from_date({2014, 4, 6})], + [2, 1, 2, "The Cap Table", days_from_date({2014, 4, 13})], + [2, 1, 3, "Articles of Incorporation", days_from_date({2014, 4, 20})], + [2, 1, 4, "Fiduciary Duties", days_from_date({2014, 4, 27})], + [2, 1, 5, "Signaling Risk", days_from_date({2014, 5, 4})], + [2, 3, 1, "Founder Friendly", days_from_date({2016, 4, 24})], + [2, 3, 2, "Two in the Box", days_from_date({2016, 5, 1})], + [2, 3, 3, "Meinertzhagen's Haversack", days_from_date({2016, 5, 8})], + [2, 3, 4, "Maleant Data Systems Solutions", days_from_date({2016, 5, 15})], + [2, 5, 1, "Grow Fast or Die Slow", days_from_date({2018, 3, 25})], + [2, 5, 2, "Reorientation", days_from_date({2018, 4, 1})], + [2, 5, 3, "Chief Operating Officer", days_from_date({2018, 4, 8})], + [2, 5, 4, "Tech Evangelist", days_from_date({2018, 4, 15})], + [2, 5, 5, "Facial Recognition", days_from_date({2018, 4, 22})], + [2, 6, 1, "Artificial Emotional Intelligence", days_from_date({2019, 10, 27})], + [2, 6, 2, "Blood Money", days_from_date({2019, 11, 3})], + [2, 6, 3, "Hooli Smokes!", days_from_date({2019, 11, 10})], + [2, 6, 4, "Maximizing Alphaness", days_from_date({2019, 11, 17})], + [2, 6, 5, "Tethics", days_from_date({2019, 11, 24})], + [2, 6, 6, "RussFest", days_from_date({2019, 12, 1})], + [2, 6, 7, "Exit Event", days_from_date({2019, 12, 8})] + ]. + +days_from_date({Year, Month, Day}) -> + calendar:date_to_gregorian_days(Year, Month, Day) - calendar:date_to_gregorian_days(1970, 1, 1). diff --git a/odbc/examples/erlang_client/src/ydb_series_client.erl b/odbc/examples/erlang_client/src/ydb_series_client.erl new file mode 100644 index 00000000000..ba3fd1cf71d --- /dev/null +++ b/odbc/examples/erlang_client/src/ydb_series_client.erl @@ -0,0 +1,221 @@ +-module(ydb_series_client). +-export([run/0, run/1, run_with_dsn/1]). + +run() -> + ConnectionString = "Driver=YDB;Endpoint=localhost:2136;Database=/local;", + run(ConnectionString). + +run(ConnectionString) when is_list(ConnectionString) -> + io:format("=== ODBC YDB Series Example ===~n"), + + application:load(odbc), + application:start(odbc), + + case odbc:connect(ConnectionString, []) of + {ok, Ref} -> + Result = run_example(Ref), + odbc:disconnect(Ref), + Result; + {error, Reason} -> + io:format("Connection failed: ~p~n", [Reason]), + error + end. + +run_with_dsn(DSN) -> + ConnectionString = lists:flatten(io_lib:format("DSN=~s;", [DSN])), + run(ConnectionString). + +run_example(Ref) -> + try + drop_tables(Ref), + create_tables(Ref), + fill_table_data(Ref), + select_simple(Ref), + upsert_simple(Ref), + select_with_params(Ref), + multistep(Ref), + select_seasons_by_series(Ref), + drop_tables(Ref), + + io:format("Completed successfully~n"), + ok + catch + Class:Reason:Stacktrace -> + io:format("~nError: ~p:~p~n", [Class, Reason]), + io:format("Stacktrace: ~p~n", [Stacktrace]), + error + end. + +create_tables(Ref) -> + Tables = [ + {"CREATE TABLE series ( + series_id Uint64, + title Utf8, + series_info Utf8, + release_date Uint64, + PRIMARY KEY (series_id) + );"}, + {"CREATE TABLE seasons ( + series_id Uint64, + season_id Uint64, + title Utf8, + first_aired Uint64, + last_aired Uint64, + PRIMARY KEY (series_id, season_id) + );"}, + {"CREATE TABLE episodes ( + series_id Uint64, + season_id Uint64, + episode_id Uint64, + title Utf8, + air_date Uint64, + PRIMARY KEY (series_id, season_id, episode_id) + );"} + ], + + lists:foreach(fun({Query}) -> + execute_update(Ref, Query) + end, Tables). + +fill_table_data(Ref) -> + SeriesData = sample_data:series(), + SeasonsData = sample_data:seasons(), + EpisodesData = sample_data:episodes(), + + batch_update(Ref, + "UPSERT INTO series (series_id, title, series_info, release_date) " + "VALUES (CAST(? AS Uint64), ?, ?, CAST(? AS Uint64))", + [{sql_integer, column(1, SeriesData)}, + {{sql_varchar, 64}, column(2, SeriesData)}, + {{sql_varchar, 256}, column(3, SeriesData)}, + {sql_integer, column(4, SeriesData)}]), + batch_update(Ref, + "UPSERT INTO seasons (series_id, season_id, title, first_aired, last_aired) " + "VALUES (CAST(? AS Uint64), CAST(? AS Uint64), ?, CAST(? AS Uint64), CAST(? AS Uint64))", + [{sql_integer, column(1, SeasonsData)}, {sql_integer, column(2, SeasonsData)}, + {{sql_varchar, 64}, column(3, SeasonsData)}, {sql_integer, column(4, SeasonsData)}, + {sql_integer, column(5, SeasonsData)}]), + batch_update(Ref, + "UPSERT INTO episodes (series_id, season_id, episode_id, title, air_date) " + "VALUES (CAST(? AS Uint64), CAST(? AS Uint64), CAST(? AS Uint64), ?, CAST(? AS Uint64))", + [{sql_integer, column(1, EpisodesData)}, {sql_integer, column(2, EpisodesData)}, + {sql_integer, column(3, EpisodesData)}, {{sql_varchar, 128}, column(4, EpisodesData)}, + {sql_integer, column(5, EpisodesData)}]), + + io:format("Inserted ~p series, ~p seasons, ~p episodes~n", + [length(SeriesData), length(SeasonsData), length(EpisodesData)]). + +select_simple(Ref) -> + Query = "SELECT series_id, title, CAST(release_date AS Date) AS release_date FROM series WHERE series_id = 1;", + Rows = selected_rows(Ref, select_simple, Query), + lists:foreach(fun(Row) -> + [Id, Title, ReleaseDate] = row_values(Row), + io:format("Series: Id=~p, Title=~p, Release=~p~n", [Id, Title, ReleaseDate]) + end, Rows). + +upsert_simple(Ref) -> + Query = "UPSERT INTO episodes (series_id, season_id, episode_id, title) VALUES (2, 6, 1, \"TBD\");", + execute_update(Ref, Query). + +select_with_params(Ref) -> + SeriesId = 2, + SeasonId = 3, + + Query = + "SELECT sa.title AS season_title, sr.title AS series_title " + "FROM seasons AS sa INNER JOIN series AS sr ON sa.series_id = sr.series_id " + "WHERE sa.series_id = CAST(? AS Uint64) AND sa.season_id = CAST(? AS Uint64);", + Params = [{sql_integer, [SeriesId]}, {sql_integer, [SeasonId]}], + + Rows = selected_param_rows(Ref, select_with_params, Query, Params), + lists:foreach(fun(Row) -> + [SeasonTitle, SeriesTitle] = row_values(Row), + io:format("Season: ~p (Series: ~p)~n", [SeasonTitle, SeriesTitle]) + end, Rows). + +multistep(Ref) -> + SeriesId = 2, + SeasonId = 5, + + Query1 = io_lib:format( + "SELECT first_aired FROM seasons WHERE series_id = ~p AND season_id = ~p;", + [SeriesId, SeasonId] + ), + + [FirstAiredRow] = selected_rows(Ref, multistep_step1, Query1), + [Date] = row_values(FirstAiredRow), + FromDate = list_to_integer(Date), + + ToDate = FromDate + 15, + + Query2 = io_lib:format( + "SELECT season_id, episode_id, title, air_date FROM episodes " + "WHERE series_id = ~p AND air_date >= ~p AND air_date <= ~p;", + [SeriesId, FromDate, ToDate] + ), + + Rows = selected_rows(Ref, multistep_step2, Query2), + lists:foreach(fun(Row) -> + [SId, EId, Title, AirDate] = row_values(Row), + io:format("Episode: S~pE~p ~p (aired: ~p)~n", [SId, EId, Title, AirDate]) + end, Rows). + +select_seasons_by_series(Ref) -> + SeriesList = [1, 2], + InClause = string:join([integer_to_list(X) || X <- SeriesList], ", "), + + Query = io_lib:format( + "SELECT series_id, season_id, title, CAST(first_aired AS Date) AS first_aired " + "FROM seasons WHERE series_id IN (~s) ORDER BY season_id;", + [InClause] + ), + + Rows = selected_rows(Ref, select_seasons_by_series, Query), + lists:foreach(fun(Row) -> + [SeriesId, SeasonId, Title, FirstAired] = row_values(Row), + io:format("Season: Series=~p, Season=~p, Title=~p, FirstAired=~p~n", + [SeriesId, SeasonId, Title, FirstAired]) + end, Rows). + +drop_tables(Ref) -> + Tables = ["series", "seasons", "episodes"], + + lists:foreach(fun(Table) -> + Query = io_lib:format("DROP TABLE ~s;", [Table]), + case odbc:sql_query(Ref, lists:flatten(Query)) of + {updated, _} -> ok; + {error, _} -> ok + end + end, Tables). + +execute_update(Ref, Query) -> + case odbc:sql_query(Ref, lists:flatten(Query)) of + {updated, _} -> ok; + Error -> throw({query_failed, update, Error}) + end. + +batch_update(Ref, Query, Params) -> + case odbc:param_query(Ref, Query, Params) of + {updated, _} -> ok; + Error -> throw({query_failed, batch_update, Error}) + end. + +column(Number, Rows) -> + [lists:nth(Number, Row) || Row <- Rows]. + +selected_rows(Ref, Step, Query) -> + case odbc:sql_query(Ref, lists:flatten(Query)) of + {selected, _, Rows} -> Rows; + Error -> throw({query_failed, Step, Error}) + end. + +selected_param_rows(Ref, Step, Query, Params) -> + case odbc:param_query(Ref, lists:flatten(Query), Params) of + {selected, _, Rows} -> Rows; + Error -> throw({query_failed, Step, Error}) + end. + +row_values(Row) when is_tuple(Row) -> + tuple_to_list(Row); +row_values(Row) -> + Row. diff --git a/odbc/examples/scheme/CMakeLists.txt b/odbc/examples/scheme/CMakeLists.txt new file mode 100644 index 00000000000..ffab881aed5 --- /dev/null +++ b/odbc/examples/scheme/CMakeLists.txt @@ -0,0 +1,14 @@ +add_executable(odbc_scheme + main.cpp +) + +target_link_libraries(odbc_scheme + PRIVATE + ODBC::ODBC +) +target_compile_definitions(odbc_scheme + PRIVATE + ODBC_DRIVER_PATH="$" +) + +add_dependencies(odbc_scheme ydb-odbc) diff --git a/odbc/examples/scheme/main.cpp b/odbc/examples/scheme/main.cpp new file mode 100644 index 00000000000..3ae2cd6fe40 --- /dev/null +++ b/odbc/examples/scheme/main.cpp @@ -0,0 +1,116 @@ +#include +#include + +#include + +void PrintOdbcError(SQLSMALLINT handleType, SQLHANDLE handle) { + SQLCHAR sqlState[6] = {0}; + SQLINTEGER nativeError = 0; + SQLCHAR message[256] = {0}; + SQLSMALLINT textLength = 0; + SQLGetDiagRec(handleType, handle, 1, sqlState, &nativeError, message, sizeof(message), &textLength); + std::cerr << "ODBC error: [" << sqlState << "] " << message << std::endl; +} + +int main() { + SQLHENV henv = nullptr; + SQLHDBC hdbc = nullptr; + SQLHSTMT hstmt = nullptr; + SQLRETURN ret; + + std::cout << "1. Allocating environment handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &henv); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating environment handle" << std::endl; + return 1; + } + SQLSetEnvAttr(henv, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0); + + std::cout << "2. Allocating connection handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_DBC, henv, &hdbc); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating connection handle" << std::endl; + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "3. Building connection string" << std::endl; + std::string connStr = "Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;"; + SQLCHAR outConnStr[1024] = {0}; + SQLSMALLINT outConnStrLen = 0; + + std::cout << "4. Connecting with SQLDriverConnect" << std::endl; + ret = SQLDriverConnect(hdbc, NULL, (SQLCHAR*)connStr.c_str(), SQL_NTS, + outConnStr, sizeof(outConnStr), &outConnStrLen, SQL_DRIVER_COMPLETE); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error connecting with SQLDriverConnect" << std::endl; + PrintOdbcError(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "5. Allocating statement handle" << std::endl; + ret = SQLAllocHandle(SQL_HANDLE_STMT, hdbc, &hstmt); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error allocating statement handle" << std::endl; + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "6. Getting tables" << std::endl; + + SQLCHAR pattern[] = "/local"; + SQLCHAR tableType[] = "TABLE"; + + ret = SQLTables(hstmt, NULL, 0, NULL, 0, pattern, SQL_NTS, tableType, SQL_NTS); + if (ret != SQL_SUCCESS && ret != SQL_SUCCESS_WITH_INFO) { + std::cerr << "Error executing query" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + SQLFreeHandle(SQL_HANDLE_STMT, hstmt); + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + return 1; + } + + std::cout << "7. Fetching result" << std::endl; + + SQLLEN ind = 0; + SQLCHAR value1[1024] = {0}; + if (SQLBindCol(hstmt, 3, SQL_C_CHAR, &value1, 1024, &ind) != SQL_SUCCESS) { + std::cerr << "Error binding column 1" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + SQLCHAR value2[1024] = {0}; + if (SQLBindCol(hstmt, 4, SQL_C_CHAR, &value2, 1024, &ind) != SQL_SUCCESS) { + std::cerr << "Error binding column 2" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + while ((ret = SQLFetch(hstmt)) == SQL_SUCCESS || ret == SQL_SUCCESS_WITH_INFO) { + if (ret != SQL_SUCCESS) { + std::cerr << "Error fetching result" << std::endl; + PrintOdbcError(SQL_HANDLE_STMT, hstmt); + return 1; + } + + std::cout << "Table name: " << value1 << std::endl; + std::cout << "Table type: " << value2 << std::endl; + + std::cout << "--------------------------------" << std::endl; + } + + std::cout << "8. Cleaning up" << std::endl; + SQLFreeHandle(SQL_HANDLE_STMT, hstmt); + SQLDisconnect(hdbc); + SQLFreeHandle(SQL_HANDLE_DBC, hdbc); + SQLFreeHandle(SQL_HANDLE_ENV, henv); + + return 0; +} diff --git a/odbc/odbc.ini b/odbc/odbc.ini new file mode 100644 index 00000000000..f7334b046f3 --- /dev/null +++ b/odbc/odbc.ini @@ -0,0 +1,9 @@ +[ODBC Data Sources] +YDB=YDB ODBC Driver + +[YDB] +Driver=YDB +Description=YDB Database Connection +Server=localhost:2136 +Database=/local +AuthMode=Anonymous diff --git a/odbc/odbcinst.ini b/odbc/odbcinst.ini new file mode 100644 index 00000000000..db2a9b8378e --- /dev/null +++ b/odbc/odbcinst.ini @@ -0,0 +1,4 @@ +[YDB] +Description=YDB ODBC Driver +Driver=/app/build/odbc/libydb-odbc.so +Setup=/app/build/odbc/libydb-odbc.so diff --git a/odbc/odbcinst.ini.in b/odbc/odbcinst.ini.in new file mode 100644 index 00000000000..8543d9adf73 --- /dev/null +++ b/odbc/odbcinst.ini.in @@ -0,0 +1,4 @@ +[YDB] +Description=YDB ODBC Driver +Driver=@YDB_ODBC_DRIVER_PATH@ +Setup=@YDB_ODBC_DRIVER_PATH@ diff --git a/odbc/packaging/postinst.in b/odbc/packaging/postinst.in new file mode 100644 index 00000000000..902817b6424 --- /dev/null +++ b/odbc/packaging/postinst.in @@ -0,0 +1,10 @@ +#!/bin/sh +set -e + +case "${1:-}" in + configure|abort-upgrade|abort-remove|abort-deconfigure) + odbcinst -i -d -f "@YDB_ODBC_DRIVER_TEMPLATE_PATH@" + ;; +esac + +exit 0 diff --git a/odbc/packaging/prerm.in b/odbc/packaging/prerm.in new file mode 100644 index 00000000000..a91d015be1b --- /dev/null +++ b/odbc/packaging/prerm.in @@ -0,0 +1,13 @@ +#!/bin/sh +set -e + +case "${1:-}" in + remove|upgrade|deconfigure) + if odbcinst -q -d -n YDB 2>/dev/null \ + | grep -Fx "Driver=@YDB_ODBC_DRIVER_PATH@" >/dev/null; then + odbcinst -u -d -n YDB + fi + ;; +esac + +exit 0 diff --git a/odbc/proposal.md b/odbc/proposal.md new file mode 100644 index 00000000000..8d98c165523 --- /dev/null +++ b/odbc/proposal.md @@ -0,0 +1,210 @@ +# ODBC driver bindings + +## Goal + +Validate the YDB ODBC driver through established ODBC bindings. Applications use the binding's public API and select YDB through a DSN or connection string. Binding implementation code is not patched. + +Languages with a maintained native YDB SDK are out of scope. PHP remains in scope because its native SDK is planned for deprecation. + +## Binding matrix + +Each language has one binding, one pinned upstream revision, its upstream database tests, and one runnable example. + +| Tier | Language | Binding | Tests | +|---|---|---|---| +| Core | Erlang | [OTP `odbc`](https://github.com/erlang/otp/tree/master/lib/odbc) | `lib/odbc/test` Common Test cases | +| Core | PHP | [PDO_ODBC](https://github.com/php/php-src/tree/master/ext/pdo_odbc) | PDO_ODBC and generic PDO PHPT tests | +| Core | Haskell | [HDBC-odbc](https://github.com/hdbc/HDBC-odbc) | HDBC/HUnit database tests | +| Core | Ruby | [ruby-odbc](https://github.com/larskanis/ruby-odbc) | Upstream test scripts | +| Core | Lua | [LuaSQL ODBC](https://github.com/lunarmodules/luasql) | Common LuaSQL and ODBC parameter tests | +| Core | Perl | [DBD::ODBC](https://github.com/perl5-dbi/DBD-ODBC) | Upstream TAP tests | +| Core | R | [odbc](https://github.com/r-dbi/odbc) | `testthat` and DBItest | +| Core | Julia | [ODBC.jl](https://github.com/JuliaDatabases/ODBC.jl) | ODBC.jl, DBInterface, and Tables tests | +| Core | Tcl | [tdbc::odbc](https://core.tcl-lang.org/tdbcodbc/timeline) | ODBC backend `tcltest` suite | +| Expansion | Raku | [DBDish::ODBC](https://github.com/salortiz/DBDish-ODBC) | Upstream and DBIish tests | +| Expansion | Crystal | [crystal-odbc](https://github.com/naqvis/crystal-odbc) | `crystal spec` | +| Expansion | Dart | [dart_odbc](https://pub.dev/packages/dart_odbc) | `dart test` | +| Expansion | D | [odbc](https://github.com/singingbush/odbc) | Upstream unit and integration tests | +| Expansion | OCaml | [ocaml-odbc](https://opam.ocaml.org/packages/odbc/) | Upstream tests | +| Expansion | Common Lisp | [CLSQL ODBC](https://github.com/sharplispers/clsql) | ODBC ASDF tests | +| Expansion | COBOL | [GixSQL ODBC](https://github.com/mridoni/gixsql) | ODBC regression tests | +| Expansion | Pascal | [Free Pascal SQLDB ODBC](https://gitlab.com/freepascal.org/fpc/source/-/tree/main/packages/fcl-db) | SQLDB connector tests | +| Expansion | Smalltalk | [Pharo-ODBC](https://github.com/pharo-rdbms/Pharo-ODBC) | SUnit tests | +| Expansion | Fortran | [odbc.f](https://davidpfister.github.io/odbc.f/) | Upstream fpm tests | + +## Required driver behavior + +### Connections + +An ODBC connection is one endpoint/database pair. YDB requires both values for routing ([connection parameters](https://ydb.tech/docs/en/concepts/connect)). + +Each `SQLHDBC` owns: + +- one SDK `TDriver` configured with its endpoint and database; +- its query, table, and scheme clients; +- its query session and active transaction; +- its current catalog, credentials, and diagnostics. + +Connection strings and DSNs must accept: + +- `Server` or `Endpoint`, `Database`, and `DSN`; +- `AuthMode=Anonymous`; +- `AuthMode=Token` with `Token`; +- `AuthMode=Static` with `User` and `Password`; +- `AuthMode=Metadata`, optionally with `MetadataHost` and `MetadataPort`; +- `AuthMode=ServiceAccount` with `ServiceAccountKeyFile`; +- `AuthMode=OAuth2` with `OAuth2KeyFile`; +- `AuthMode=Environment`; +- `IamEndpoint`, `RootCertificate`, `ClientCertificate`, and `ClientPrivateKey`. + +`UID`/`PWD`, `AccessToken`, `SaFile`, and `CaFile` are accepted aliases. The +authentication mode is inferred when exactly one credential type is present. +`SQLConnect` user and password arguments override DSN values. + +The current implementation follows this model: `SQLAllocHandle(SQL_HANDLE_DBC)` creates a `TConnection`, and `TConnection::TYdbState` owns the SDK driver and clients. `SQLHENV` only tracks connection handles. Connections in the same environment therefore keep endpoint, database, session, transaction, and catalog state separate. + +`SQL_ATTR_CURRENT_CATALOG` uses a path below the connected database as `TablePathPrefix`. Setting it to another database path recreates only that connection's SDK state. + +Required tests: + +- two database paths on one endpoint; +- two endpoints; +- simultaneous queries; +- independent transactions and catalogs; +- failure and disconnect of one connection while the other remains usable; +- driver-manager pooling keyed by endpoint, database, credentials, and TLS settings. + +### Cursors + +YDB returns result sets, not server-side ODBC cursors. The driver owns cursor state for each statement. + +The cursor states are `before first`, `on row or rowset`, `after last`, and `closed`. `SQLFetch` and `SQLFetchScroll(SQL_FETCH_NEXT)` advance a forward cursor. A static cursor stores a result snapshot and implements `FIRST`, `LAST`, `PRIOR`, `ABSOLUTE`, and `RELATIVE`. + +The cursor also owns: + +- typed YDB rows and column metadata; +- row-wise and column-wise bindings; +- row-array status and processed-row counters; +- a separate chunk offset for each `SQLGetData` column; +- bounded memory with a statement-local spill file for static cursors. + +Re-execution replaces the cursor. Close, cancel, commit, rollback, and disconnect release it according to the advertised cursor behavior. + +### Binding contract + +The driver must support the operations used by the Core bindings: + +- DSN and connection-string connection; +- prepare, bind, execute, and data-at-execution parameters; +- scalar and rowset fetch; +- `NULL`, integer, floating-point, decimal, text, binary, date, time, and timestamp conversion; +- column, table, key, index, type, and result metadata; +- autocommit and explicit transactions; +- diagnostics through SQLSTATE and native YDB issues; +- independent statements and connections; +- deterministic cleanup after errors. + +### Row counts + +Data-modification statements must request Basic YDB query statistics. +`SQLRowCount` must sum updated and deleted rows from every query phase. +Parameter-array execution must sum the count of each executed parameter set. It +returns `-1` for other statements or when statistics do not contain a usable +count. + +### Debian package + +The driver is shipped as a separate `ydb-odbc` package with the same version as +the SDK release. It contains `libydb-odbc.so` in the multiarch library directory +and an unixODBC driver template. Package dependencies include `odbcinst` and the +shared-library dependencies derived from the built artifact. + +Installation registers the `YDB` driver with `odbcinst -i -d -f`. Upgrade +updates the registration without creating duplicate entries. Removal +unregisters only the entry owned by the package. The package does not install +`/etc/odbc.ini` or modify user DSNs. + +The package is built and published with the SDK release. A clean-container test +installs it, checks `odbcinst -q -d`, connects through `isql` and Qt QODBC, +tests an upgrade, removes the package, and verifies that unrelated drivers and +user DSNs remain unchanged. + +## Current limitations + +| Area | Limitation | +|---|---| +| SQL dialect | Statements are YQL. The driver rewrites ODBC escapes and `?` parameters; it is not a general ANSI SQL translator. | +| Authentication | Connection strings currently configure only endpoint, database, and DSN. Authentication and TLS settings are not wired into `TDriverConfig`. | +| Retry classification | The driver cannot infer whether arbitrary SQL is idempotent. Autocommit statement retries use `TRetryOperationSettings::Idempotent(false)`. This enables only retries safe for a non-idempotent operation and may return an error with an unknown execution outcome. See [YDB retry settings](https://ydb.tech/docs/en/recipes/ydb-sdk/retry) and [error handling](https://ydb.tech/docs/en/reference/ydb-sdk/error_handling). | +| Explicit transactions | An ODBC transaction is not retried by the driver. Conflicts, node failures, maintenance, and network failures can abort it, including at commit. The application must open a new transaction and replay the entire unit of work in a retry loop. Retrying only the failed statement is incorrect ([query execution](https://ydb.tech/docs/en/concepts/query_execution/), [transactions](https://ydb.tech/docs/en/concepts/transactions)). | +| Transaction isolation | Read-write connections support serializable and snapshot read-write modes. ODBC read-committed and read-uncommitted requests are rejected. | +| Transaction scope | `SQLEndTran(SQL_HANDLE_ENV, ...)` completes each connection independently. It is not an atomic transaction across databases or endpoints. | +| Connection concurrency | An explicit transaction uses one SDK session. YDB sessions execute one query at a time, so concurrent statements on the same transaction connection require application serialization ([YDB errors FAQ](https://ydb.tech/docs/en/faq/errors)). | +| Cursor support | The current implementation is forward-only. `SQLFetchScroll` accepts only `SQL_FETCH_NEXT`; static scrolling and spill are planned. | +| Result sets | Only the first result set is exposed. `SQLMoreResults` returns `SQL_NO_DATA`. | +| Result buffering | Query execution uses the non-streaming SDK result and keeps it for cursor fetches. Large results can consume memory proportional to the result size. | +| Row counts | `SQLRowCount` currently returns `-1`; affected-row counts are not extracted from YDB query statistics. | +| Prepare | `SQLPrepare` stores the query and counts client-side parameter markers. It does not create a persistent server-side prepared statement. | +| Batches | Parameter arrays are accepted only for data-modification statements and execute sequentially. Earlier parameter sets may already be committed when a later set fails. | +| Cancellation | Execution is synchronous. `SQLCancel` clears local cursor and parameter state but does not interrupt an in-flight SDK request. | +| Metadata namespace | YDB paths are exposed as catalogs. Schemas are empty. | +| DDL | Autocommit DDL uses `NoTx`. DDL executed while autocommit is off is sent through the active transaction and may be rejected by YDB. | +| Optional ODBC features | Multiple result sets, stored procedures, output parameters, positioned updates, bookmarks, asynchronous execution, and ODBC batch operations are not implemented. | +| Thread safety | Handle state is mutable and has no internal locking. Applications must serialize access to the same ODBC handle. | + +## Test repository + +```text +odbc/tests/frameworks/ + registry.yaml + / + upstream.lock + run-tests + convert-results + example/ +odbc/tests/reporting/ +``` + +`registry.yaml` records the language, tier, runtime image, upstream URL and revision, archive checksum, test command, result format, and example command. CI verifies the checksum before running the unchanged upstream binding. + +Each example uses only the binding's public API. It accepts `YDB_ODBC_DSN` or `YDB_ODBC_CONNECTION_STRING`, creates isolated test data, executes bound statements, iterates a cursor, demonstrates commit and rollback, and cleans up. + +Test output is converted to Allure without changing the upstream runner. Reports include the binding version, runtime version, driver commit, YDB version, endpoint/database mode, and upstream checksum. Missing tests, an empty suite, infrastructure failure, and unexpected skips fail the job. + +## CI + +The complete matrix runs only: + +- after a merge into `odbc-driver-feature`; +- when a pull request has the special full-matrix tag. + +It does not run for ordinary pull requests, direct non-merge pushes, Git tags, schedules, or manual dispatches. + +Each job starts a pinned YDB version, installs the `ydb-odbc` package, creates a +job-local DSN, verifies the upstream source, runs the upstream tests and +example, and uploads native and Allure results. + +## Delivery order + +1. Stabilize the driver build and add the `ydb-odbc` package component. +2. Add clean install, upgrade, removal, unixODBC registration, `isql`, and Qt + QODBC package tests. +3. Complete connection settings, authentication, TLS, row counts, and + endpoint/database isolation. +4. Add static cursor emulation and cursor integration tests. +5. Add the registry, framework runner template, source verification, result + conversion, and report aggregation. +6. Onboard Core bindings and enable the full-matrix gate. +7. Onboard Expansion bindings after the initial release. + +## Acceptance + +- `ydb-odbc` passes clean install, upgrade, removal, registration, `isql`, and + Qt QODBC tests. +- Every selected binding runs from a pinned, verified upstream source. +- Binding implementation code is unchanged. +- Every upstream database test is executed or reported as unsupported with its original test identifier. +- Every binding has a runnable example. +- Unit and integration tests pass. +- Endpoint/database isolation tests pass. +- Reports contain no missing or unexpected skipped tests. diff --git a/odbc/src/connection.cpp b/odbc/src/connection.cpp new file mode 100644 index 00000000000..b055b965d2d --- /dev/null +++ b/odbc/src/connection.cpp @@ -0,0 +1,338 @@ +#include "connection.h" +#include "statement.h" + +#include +#include + +#include +#include +#include + +#include +#include + +namespace NYdb::NOdbc { + +TConnection::~TConnection() { + DestroyYdbState(); +} + +void TConnection::DestroyYdbState() { + QuerySession_.reset(); + Tx_.reset(); + Ydb_.reset(); +} + +SQLRETURN TConnection::DriverConnect(std::string_view connectionString) { + std::vector ignoredAttributes; + TConnectionParameters explicitParameters = + ParseAndNormalizeConnectionString(connectionString, ignoredAttributes); + const auto dsnIt = explicitParameters.find("DSN"); + TConnectionParameters parameters; + if (dsnIt != explicitParameters.end() && !dsnIt->second.empty()) { + parameters = ReadDsnParameters(dsnIt->second); + } + OverlayConnectionParameters(parameters, explicitParameters); + ApplyResolvedSettings(ResolveConnectionSettings(std::move(parameters))); + + if (!ignoredAttributes.empty()) { + std::string message = ignoredAttributes.size() == 1 + ? "Invalid connection string attribute ignored: " + : "Invalid connection string attributes ignored: "; + for (size_t i = 0; i < ignoredAttributes.size(); ++i) { + if (i != 0) { + message += ", "; + } + message += ignoredAttributes[i]; + } + return AddError("01S00", 0, message, SQL_SUCCESS_WITH_INFO); + } + + return SQL_SUCCESS; +} + +SQLRETURN TConnection::Connect(std::string_view serverName, + std::string_view userName, + std::string_view auth) { + TConnectionParameters parameters = ReadDsnParameters(serverName); + if (!userName.empty() || !auth.empty()) { + for (const std::string_view key : { + "Token", "MetadataHost", "MetadataPort", "ServiceAccountKeyFile", + "OAuth2KeyFile", "IamEndpoint"}) { + parameters.erase(std::string(key)); + } + parameters["AuthMode"] = "Static"; + } + if (!userName.empty()) { + parameters["User"] = std::string(userName); + } + if (!auth.empty()) { + parameters["Password"] = std::string(auth); + } + ApplyResolvedSettings(ResolveConnectionSettings(std::move(parameters), std::string(serverName))); + + return SQL_SUCCESS; +} + +SQLRETURN TConnection::Disconnect() { + DestroyYdbState(); + DriverConfig_.reset(); + DbmsVersionCache_.reset(); + Endpoint_.clear(); + Database_.clear(); + DataSourceName_.clear(); + return SQL_SUCCESS; +} + +NQuery::TSession& TConnection::GetOrCreateQuerySession() { + if (!QuerySession_) { + auto sessionResult = Ydb_->QueryClient.GetSession().ExtractValueSync(); + NStatusHelpers::ThrowOnError(sessionResult); + QuerySession_.emplace(std::move(sessionResult.GetSession())); + } + return *QuerySession_; +} + +std::optional TConnection::GetClient() { + if (!Ydb_) { + return std::nullopt; + } + return Ydb_->QueryClient; +} + +std::optional TConnection::GetTableClient() { + if (!Ydb_) { + return std::nullopt; + } + return Ydb_->TableClient; +} + +std::optional TConnection::GetSchemeClient() { + if (!Ydb_) { + return std::nullopt; + } + return Ydb_->SchemeClient; +} + +std::unique_ptr TConnection::CreateStatement() { + return std::make_unique(this); +} + +void TConnection::CloseStatementCursors() { + for (TStatement* stmt : Statements_) { + stmt->Close(true); + } +} + +SQLRETURN TConnection::SetAutocommit(bool value) { + if (value && Tx_) { + auto status = Tx_->Commit().ExtractValueSync(); + NStatusHelpers::ThrowOnError(status); + Tx_.reset(); + } + return Attributes_.SetAutocommit(value); +} + +bool TConnection::GetAutocommit() const { + return Attributes_.GetAutocommit(); +} + +SQLRETURN TConnection::SetConnectAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER stringLength) { + if (attr == SQL_ATTR_CURRENT_CATALOG) { + std::optional rebindDatabase; + SQLRETURN rc = Attributes_.ApplyCatalogChange(value, stringLength, Database_, rebindDatabase, *this); + if (rc != SQL_SUCCESS) { + return rc; + } + if (rebindDatabase) { + RebindToDatabase(*rebindDatabase); + } + return SQL_SUCCESS; + } + return Attributes_.SetConnectAttr(attr, value, stringLength, [this](bool autocommit) { + return SetAutocommit(autocommit); + }, *this); +} + +SQLRETURN TConnection::GetConnectAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr) { + return Attributes_.GetConnectAttr(attr, value, bufferLength, stringLengthPtr, *this); +} + +NQuery::TTxSettings TConnection::MakeTxSettings() const { + return Attributes_.MakeTxSettings(); +} + +const std::optional& TConnection::GetTx() { + return Tx_; +} + +void TConnection::SetTx(const NQuery::TTransaction& tx) { + Tx_ = tx; +} + +void TConnection::ResetTx() { + Tx_.reset(); +} + +void TConnection::ResetQuerySession() { + QuerySession_.reset(); +} + +SQLRETURN TConnection::CommitTx() { + if (!Tx_) { + return SQL_SUCCESS; + } + auto status = Tx_->Commit().ExtractValueSync(); + NStatusHelpers::ThrowOnError(status); + Tx_.reset(); + CloseStatementCursors(); + return SQL_SUCCESS; +} + +SQLRETURN TConnection::RollbackTx() { + if (!Tx_) { + return SQL_SUCCESS; + } + auto status = Tx_->Rollback().ExtractValueSync(); + NStatusHelpers::ThrowOnError(status); + Tx_.reset(); + CloseStatementCursors(); + return SQL_SUCCESS; +} + +void TConnection::SetEnvironment(TEnvironment* env){ + if (ParentEnv_){ + throw std::logic_error("Connection already bound to environment"); + } + ParentEnv_ = env; +} + +TEnvironment* TConnection::GetEnvironment(){ + return ParentEnv_; +} + +const std::string& TConnection::GetDataSourceName() const { + return DataSourceName_; +} + +SQLUINTEGER TConnection::GetSupportedTxnIsolationOptions() const { + return Attributes_.GetSupportedTxnIsolationOptions(); +} + +bool TConnection::IsDataSourceReadOnly() const { + return Attributes_.GetAccessMode() == SQL_MODE_READ_ONLY; +} + +const std::string& TConnection::GetDbmsVersion() { + if (DbmsVersionCache_) { + return *DbmsVersionCache_; + } + + auto client = GetClient(); + if (!client) { + throw TOdbcException("08003", 0, "Connection is not established"); + } + + std::optional fetched; + const NYdb::TStatus status = client->RetryQuerySync( + [&fetched](NQuery::TSession session) -> NYdb::TStatus { + auto result = session.ExecuteQuery( + "SELECT Version();", + NQuery::TTxControl::NoTx(), + NYdb::TParamsBuilder().Build()).ExtractValueSync(); + if (!result.IsSuccess()) { + return result; + } + if (result.GetResultSets().empty()) { + return NYdb::TStatus(EStatus::SUCCESS, NYdb::NIssue::TIssues()); + } + TResultSetParser parser(result.GetResultSetParser(0)); + if (parser.TryNextRow()) { + fetched = parser.ColumnParser(0).GetUtf8(); + } + return NYdb::TStatus(EStatus::SUCCESS, NYdb::NIssue::TIssues()); + }); + + NStatusHelpers::ThrowOnError(status); + if (!fetched || fetched->empty()) { + throw TOdbcException("HY000", 0, "Failed to retrieve DBMS version"); + } + + DbmsVersionCache_ = std::move(*fetched); + return *DbmsVersionCache_; +} + +void TConnection::RecreateYdbClients() { + if (!DriverConfig_) { + throw TOdbcException("08003", 0, "Connection configuration is not available"); + } + DestroyYdbState(); + DbmsVersionCache_.reset(); + Ydb_.emplace(*DriverConfig_); +} + +void TConnection::ApplyResolvedSettings(TResolvedConnectionSettings&& settings) { + TConnectionAttributes::NormalizeCatalogPath(settings.Database); + settings.DriverConfig.SetDatabase(settings.Database); + + Endpoint_ = std::move(settings.Endpoint); + Database_ = std::move(settings.Database); + DataSourceName_ = std::move(settings.DataSourceName); + DriverConfig_.emplace(std::move(settings.DriverConfig)); + RecreateYdbClients(); + Attributes_.SetCurrentCatalog(Database_); +} + +void TConnection::RebindToDatabase(std::string_view newDatabase) { + if (!DriverConfig_) { + throw TOdbcException("08003", 0, "Connection configuration is not available"); + } + std::string db(newDatabase); + TConnectionAttributes::NormalizeCatalogPath(db); + Database_ = std::move(db); + DriverConfig_->SetDatabase(Database_); + Attributes_.SetCurrentCatalog(Database_); + RecreateYdbClients(); +} + + +std::string TConnection::WrapQueryForCurrentCatalog(const std::string& sql) const { + std::optional rel = Attributes_.ResolveCatalogRoute(Database_).TablePathPrefix; + if (!rel) { + return sql; + } + std::string escapedPrefix; + escapedPrefix.reserve(rel->size() + 8); + for (const char ch : *rel) { + if (ch == '\\' || ch == '"') { + escapedPrefix.push_back('\\'); + } + escapedPrefix.push_back(ch); + } + return "PRAGMA TablePathPrefix = \"" + escapedPrefix + "\";\n" + sql; +} + +SQLRETURN TConnection::NativeSql(const std::string& inSql, SQLCHAR* outSql, SQLINTEGER outMax, SQLINTEGER* outLen) { + const SQLINTEGER fullLen = static_cast(inSql.size()); + if (outLen) { + *outLen = fullLen; + } + if (!outSql) { + return outMax == 0 ? SQL_SUCCESS : AddError("HY090", 0, "Invalid string or buffer length"); + } + if (outMax <= 0) { + return fullLen == 0 ? SQL_SUCCESS : AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + } + const SQLINTEGER copyLen = std::min(fullLen, outMax - 1); + if (copyLen > 0) { + std::memcpy(outSql, inSql.data(), static_cast(copyLen)); + } + outSql[copyLen] = '\0'; + if (copyLen < fullLen) { + return AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + } + return SQL_SUCCESS; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/connection.h b/odbc/src/connection.h new file mode 100644 index 00000000000..74437b73eb5 --- /dev/null +++ b/odbc/src/connection.h @@ -0,0 +1,118 @@ +#pragma once + +#include "environment.h" +#include "connection_attr.h" +#include "connection_config.h" +#include "utils/error_manager.h" + +#include +#include +#include +#include + +#include +#include + +#include +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +class TStatement; +class TDescriptor; + +class TConnection : public TErrorManager { +private: + struct TYdbState { + // Declared first: constructed before clients, destroyed after them. + TDriver Driver; + NQuery::TQueryClient QueryClient; + NScheme::TSchemeClient SchemeClient; + NTable::TTableClient TableClient; + + explicit TYdbState(const TDriverConfig& config) + : Driver(config) + , QueryClient(Driver) + , SchemeClient(Driver) + , TableClient(Driver) + {} + + ~TYdbState() { + Driver.Stop(true); + } + }; + + std::optional Ydb_; + std::optional DriverConfig_; + std::optional Tx_; + std::optional QuerySession_; + + std::string Endpoint_; + std::string Database_; + std::string DataSourceName_; + TEnvironment* ParentEnv_ = nullptr; + + TConnectionAttributes Attributes_; + mutable std::optional DbmsVersionCache_; + std::unordered_set Statements_; + std::unordered_set Descriptors_; + + void DestroyYdbState(); + void ApplyResolvedSettings(TResolvedConnectionSettings&& settings); + void RecreateYdbClients(); + void RebindToDatabase(std::string_view newDatabase); +public: + ~TConnection(); + + SQLRETURN Connect(std::string_view serverName, + std::string_view userName, + std::string_view auth); + + SQLRETURN DriverConnect(std::string_view connectionString); + SQLRETURN Disconnect(); + + std::unique_ptr CreateStatement(); + void RegisterStatement(TStatement* stmt) { Statements_.insert(stmt); } + void UnregisterStatement(TStatement* stmt) { Statements_.erase(stmt); } + void RegisterDescriptor(TDescriptor* desc) { Descriptors_.insert(desc); } + void UnregisterDescriptor(TDescriptor* desc) { Descriptors_.erase(desc); } + bool HasChildren() const noexcept { return !Statements_.empty() || !Descriptors_.empty(); } + void CloseStatementCursors(); + + std::optional GetClient(); + NQuery::TSession& GetOrCreateQuerySession(); + std::optional GetTableClient(); + std::optional GetSchemeClient(); + + SQLRETURN SetAutocommit(bool value); + bool GetAutocommit() const; + + SQLRETURN SetConnectAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER stringLength); + SQLRETURN GetConnectAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr); + NQuery::TTxSettings MakeTxSettings() const; + + std::string WrapQueryForCurrentCatalog(const std::string& sql) const; + const std::string& GetDbmsVersion(); + const std::string& GetDataSourceName() const; + SQLUINTEGER GetSupportedTxnIsolationOptions() const; + bool IsDataSourceReadOnly() const; + + const std::optional& GetTx(); + void SetTx(const NQuery::TTransaction& tx); + void ResetTx(); + void ResetQuerySession(); + + SQLRETURN CommitTx(); + SQLRETURN RollbackTx(); + + void SetEnvironment(TEnvironment* env); + TEnvironment* GetEnvironment(); + + SQLRETURN NativeSql(const std::string& inSql, SQLCHAR* outSql, SQLINTEGER outMax, SQLINTEGER* outLen); +}; + +} // namespace NYdb::NOdbc diff --git a/odbc/src/connection_attr.cpp b/odbc/src/connection_attr.cpp new file mode 100644 index 00000000000..c1e350d5e39 --- /dev/null +++ b/odbc/src/connection_attr.cpp @@ -0,0 +1,354 @@ + +#include "connection_attr.h" +#include "utils/attr.h" +#include "utils/diag.h" + +#include + +namespace NYdb { +namespace NOdbc { + +namespace { + +namespace Catalog { + +void NormalizePath(std::string& path) { + if (path.empty() || path == "/") { + return; + } + const size_t trailingSlashStart = path.find_last_not_of('/'); + if (trailingSlashStart == std::string::npos) { + path.assign("/"); + return; + } + path.erase(trailingSlashStart + 1); +} + +TConnectionAttributes::TCatalogBinding BuildBinding(const std::string& currentCatalog, const std::string& database) { + TConnectionAttributes::TCatalogBinding binding; + binding.Catalog = currentCatalog; + binding.Database = database; + NormalizePath(binding.Catalog); + NormalizePath(binding.Database); + if (binding.Catalog == binding.Database) { + return binding; + } + + const std::string databasePrefix = binding.Database + "/"; + if (binding.Catalog.size() <= databasePrefix.size() || + binding.Catalog.compare(0, databasePrefix.size(), databasePrefix) != 0) { + return binding; + } + + std::string relativeCatalog = binding.Catalog.substr(databasePrefix.size()); + if (!relativeCatalog.empty()) { + binding.RelativeCatalog = std::move(relativeCatalog); + } + return binding; +} + +} // namespace Catalog + +namespace Tx { + +bool IsKnownTxnIsolation(SQLUINTEGER txnIsolation) { + switch (txnIsolation) { + case SQL_TXN_READ_UNCOMMITTED: + case SQL_TXN_READ_COMMITTED: + case SQL_TXN_REPEATABLE_READ: + case SQL_TXN_SERIALIZABLE: + return true; + default: + return false; + } +} + +std::optional ResolveTxMode(SQLUINTEGER accessMode, SQLUINTEGER txnIsolation) { + if (accessMode == SQL_MODE_READ_ONLY) { + return NQuery::TTxSettings::TS_SNAPSHOT_RO; + } + + switch (txnIsolation) { + case SQL_TXN_REPEATABLE_READ: + return NQuery::TTxSettings::TS_SNAPSHOT_RW; + case SQL_TXN_SERIALIZABLE: + return NQuery::TTxSettings::TS_SERIALIZABLE_RW; + default: + return std::nullopt; + } +} + +} // namespace Tx + +namespace Autocommit { + +SQLRETURN Get(bool autocommitEnabled, SQLPOINTER value) { + auto* out = reinterpret_cast(value); + *out = autocommitEnabled ? SQL_AUTOCOMMIT_ON : SQL_AUTOCOMMIT_OFF; + return SQL_SUCCESS; +} + +} // namespace Autocommit + +} + +void TConnectionAttributes::NormalizeCatalogPath(std::string& path) { + Catalog::NormalizePath(path); +} + +SQLRETURN TConnectionAttributes::SetAutocommit(bool value) { + Autocommit_ = value; + return SQL_SUCCESS; +} + +bool TConnectionAttributes::GetAutocommit() const { + return Autocommit_; +} + +SQLRETURN TConnectionAttributes::SetConnectAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER stringLength, + const std::function& applyAutocommit, + TErrorManager& errors) { + switch (attr) { + case SQL_ATTR_AUTOCOMMIT: + return SetAutocommit(value, applyAutocommit, errors); + case SQL_ATTR_ACCESS_MODE: + return SetAccessMode(value, errors); + case SQL_ATTR_TXN_ISOLATION: + return SetTxnIsolation(value, errors); + case SQL_ATTR_CURRENT_CATALOG: + return SetCurrentCatalog(value, stringLength, errors); + case SQL_ATTR_QUIET_MODE: + QuietMode_ = value; + return SQL_SUCCESS; + case SQL_ATTR_TRANSLATE_LIB: + if (!value) { + return Diag::AddNullPointer(errors); + } + return errors.AddError( + "HYC00", 0, "Translation libraries are not supported"); + case SQL_ATTR_TRANSLATE_OPTION: + TranslateOption_ = ReadIntegerAttr(value); + return SQL_SUCCESS; + default: + return Diag::AddNotImplemented(errors); + } +} + +SQLRETURN TConnectionAttributes::GetConnectAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const { + const bool stringAttribute = + attr == SQL_ATTR_CURRENT_CATALOG || attr == SQL_ATTR_TRANSLATE_LIB; + if (!value && (!stringAttribute || !stringLengthPtr)) { + return Diag::AddNullPointer(errors); + } + if (stringLengthPtr) { + *stringLengthPtr = 0; + } + switch (attr) { + case SQL_ATTR_AUTOCOMMIT: + return GetAutocommit(value); + case SQL_ATTR_ACCESS_MODE: + return GetAccessMode(value); + case SQL_ATTR_TXN_ISOLATION: + return GetTxnIsolation(value); + case SQL_ATTR_CURRENT_CATALOG: + return GetCurrentCatalog(value, bufferLength, stringLengthPtr, errors); + case SQL_ATTR_QUIET_MODE: + if (!QuietMode_) { + return SQL_NO_DATA; + } + *reinterpret_cast(value) = *QuietMode_; + if (stringLengthPtr) { + *stringLengthPtr = sizeof(SQLPOINTER); + } + return SQL_SUCCESS; + case SQL_ATTR_TRANSLATE_LIB: + return SQL_NO_DATA; + case SQL_ATTR_TRANSLATE_OPTION: + if (!TranslateOption_) { + return SQL_NO_DATA; + } + *reinterpret_cast(value) = *TranslateOption_; + if (stringLengthPtr) { + *stringLengthPtr = sizeof(SQLUINTEGER); + } + return SQL_SUCCESS; + default: + return Diag::AddNotImplemented(errors); + } +} + +SQLRETURN TConnectionAttributes::SetAutocommit( + SQLPOINTER value, + const std::function& applyAutocommit, + TErrorManager& errors) { + const auto token = ReadIntegerAttrIfIn( + value, + {static_cast(SQL_AUTOCOMMIT_ON), static_cast(SQL_AUTOCOMMIT_OFF)}); + if (!token) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_AUTOCOMMIT"); + } + if (*token == static_cast(SQL_AUTOCOMMIT_ON)) { + return applyAutocommit(true); + } + return applyAutocommit(false); +} + +SQLRETURN TConnectionAttributes::SetAccessMode(SQLPOINTER value, TErrorManager& errors) { + const auto mode = ReadIntegerAttrIfIn(value, {SQL_MODE_READ_WRITE, SQL_MODE_READ_ONLY}); + if (!mode) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_ACCESS_MODE"); + } + auto txMode = Tx::ResolveTxMode(*mode, TxnIsolation_); + if (!txMode) { + return errors.AddError( + "HYC00", + 0, + *mode == SQL_MODE_READ_WRITE + ? "Transaction isolation is not supported for read-write mode" + : "Transaction isolation is not supported for read-only mode"); + } + AccessMode_ = *mode; + TxMode_ = *txMode; + return SQL_SUCCESS; +} + +SQLRETURN TConnectionAttributes::SetTxnIsolation(SQLPOINTER value, TErrorManager& errors) { + const SQLUINTEGER isolation = ReadIntegerAttr(value); + if (!Tx::IsKnownTxnIsolation(isolation)) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_TXN_ISOLATION"); + } + auto txMode = Tx::ResolveTxMode(AccessMode_, isolation); + if (!txMode) { + return errors.AddError("HYC00", 0, "SQL_ATTR_TXN_ISOLATION value is not supported"); + } + TxnIsolation_ = isolation; + TxMode_ = *txMode; + return SQL_SUCCESS; +} + +SQLRETURN TConnectionAttributes::SetCurrentCatalog(SQLPOINTER value, SQLINTEGER stringLength, TErrorManager& errors) { + if (!value) { + return Diag::AddNullPointer(errors); + } + std::string catalog = ReadAttributeString(value, stringLength); + Catalog::NormalizePath(catalog); + if (catalog.empty()) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_CURRENT_CATALOG"); + } + CurrentCatalog_ = std::move(catalog); + return SQL_SUCCESS; +} + +SQLRETURN TConnectionAttributes::GetAutocommit(SQLPOINTER value) const { + return Autocommit::Get(Autocommit_, value); +} + +SQLRETURN TConnectionAttributes::GetAccessMode(SQLPOINTER value) const { + auto* out = reinterpret_cast(value); + *out = AccessMode_; + return SQL_SUCCESS; +} + +SQLRETURN TConnectionAttributes::GetTxnIsolation(SQLPOINTER value) const { + auto* out = reinterpret_cast(value); + *out = TxnIsolation_; + return SQL_SUCCESS; +} + +SQLUINTEGER TConnectionAttributes::GetAccessMode() const { + return AccessMode_; +} + +SQLUINTEGER TConnectionAttributes::GetSupportedTxnIsolationOptions() const { + static constexpr SQLUINTEGER kLevels[] = { + SQL_TXN_READ_UNCOMMITTED, + SQL_TXN_READ_COMMITTED, + SQL_TXN_REPEATABLE_READ, + SQL_TXN_SERIALIZABLE, + }; + SQLUINTEGER mask = 0; + for (const SQLUINTEGER level : kLevels) { + if (Tx::ResolveTxMode(AccessMode_, level)) { + mask |= level; + } + } + return mask; +} + +SQLRETURN TConnectionAttributes::GetCurrentCatalog( + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const { + return WriteAttributeString(CurrentCatalog_, value, bufferLength, stringLengthPtr, errors); +} + +NQuery::TTxSettings TConnectionAttributes::MakeTxSettings() const { + switch (TxMode_) { + case NQuery::TTxSettings::TS_ONLINE_RO: + return NQuery::TTxSettings::OnlineRO(); + case NQuery::TTxSettings::TS_STALE_RO: + return NQuery::TTxSettings::StaleRO(); + case NQuery::TTxSettings::TS_SNAPSHOT_RO: + return NQuery::TTxSettings::SnapshotRO(); + case NQuery::TTxSettings::TS_SNAPSHOT_RW: + return NQuery::TTxSettings::SnapshotRW(); + case NQuery::TTxSettings::TS_SERIALIZABLE_RW: + default: + return NQuery::TTxSettings::SerializableRW(); + } +} + +void TConnectionAttributes::SetCurrentCatalog(const std::string& value) { + CurrentCatalog_ = value; + Catalog::NormalizePath(CurrentCatalog_); +} + +const std::string& TConnectionAttributes::GetCurrentCatalog() const { + return CurrentCatalog_; +} + +TConnectionAttributes::TCatalogBinding TConnectionAttributes::BuildCatalogBinding(const std::string& database) const { + return Catalog::BuildBinding(CurrentCatalog_, database); +} + +TConnectionAttributes::TCatalogRoute TConnectionAttributes::ResolveCatalogRoute(const std::string& currentDatabase) const { + const TCatalogBinding binding = BuildCatalogBinding(currentDatabase); + if (binding.Catalog == binding.Database) { + return {binding.Database, std::nullopt}; + } + if (binding.RelativeCatalog) { + return {binding.Database, binding.Catalog}; + } + return {binding.Catalog, std::nullopt}; +} + +SQLRETURN TConnectionAttributes::ApplyCatalogChange( + SQLPOINTER value, + SQLINTEGER stringLength, + const std::string& currentDatabase, + std::optional& rebindDatabase, + TErrorManager& errors) { + SQLRETURN rc = SetCurrentCatalog(value, stringLength, errors); + if (rc != SQL_SUCCESS) { + return rc; + } + const TCatalogRoute route = ResolveCatalogRoute(currentDatabase); + if (route.EffectiveDatabase != currentDatabase) { + rebindDatabase = route.EffectiveDatabase; + } else { + rebindDatabase.reset(); + } + return SQL_SUCCESS; +} + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/connection_attr.h b/odbc/src/connection_attr.h new file mode 100644 index 00000000000..8f6a0614af4 --- /dev/null +++ b/odbc/src/connection_attr.h @@ -0,0 +1,90 @@ +#pragma once + +#include "utils/error_manager.h" + +#include + +#include +#include +#include + +#include +#include + +namespace NYdb { +namespace NOdbc { + +class TConnectionAttributes { +public: + struct TCatalogBinding { + std::string Catalog; + std::string Database; + std::optional RelativeCatalog; + }; + + struct TCatalogRoute { + std::string EffectiveDatabase; + std::optional TablePathPrefix; + }; + + SQLRETURN SetAutocommit(bool value); + bool GetAutocommit() const; + + SQLRETURN SetConnectAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER stringLength, + const std::function& applyAutocommit, + TErrorManager& errors); + + SQLRETURN GetConnectAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const; + + NQuery::TTxSettings MakeTxSettings() const; + void SetCurrentCatalog(const std::string& value); + const std::string& GetCurrentCatalog() const; + TCatalogBinding BuildCatalogBinding(const std::string& database) const; + TCatalogRoute ResolveCatalogRoute(const std::string& currentDatabase) const; + SQLRETURN ApplyCatalogChange( + SQLPOINTER value, + SQLINTEGER stringLength, + const std::string& currentDatabase, + std::optional& rebindDatabase, + TErrorManager& errors); + static void NormalizeCatalogPath(std::string& path); + SQLUINTEGER GetSupportedTxnIsolationOptions() const; + SQLUINTEGER GetAccessMode() const; + +private: + SQLRETURN SetAutocommit( + SQLPOINTER value, + const std::function& applyAutocommit, + TErrorManager& errors); + SQLRETURN SetAccessMode(SQLPOINTER value, TErrorManager& errors); + SQLRETURN SetTxnIsolation(SQLPOINTER value, TErrorManager& errors); + SQLRETURN SetCurrentCatalog(SQLPOINTER value, SQLINTEGER stringLength, TErrorManager& errors); + + SQLRETURN GetAutocommit(SQLPOINTER value) const; + SQLRETURN GetAccessMode(SQLPOINTER value) const; + SQLRETURN GetTxnIsolation(SQLPOINTER value) const; + SQLRETURN GetCurrentCatalog( + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const; + + bool Autocommit_ = true; + std::string CurrentCatalog_; + std::optional QuietMode_; + std::optional TranslateOption_; + SQLUINTEGER AccessMode_ = SQL_MODE_READ_WRITE; + SQLUINTEGER TxnIsolation_ = SQL_TXN_SERIALIZABLE; + NQuery::TTxSettings::ETransactionMode TxMode_ = NQuery::TTxSettings::TS_SERIALIZABLE_RW; +}; + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/connection_config.cpp b/odbc/src/connection_config.cpp new file mode 100644 index 00000000000..db1735def5d --- /dev/null +++ b/odbc/src/connection_config.cpp @@ -0,0 +1,435 @@ +#include "connection_config.h" + +#include "utils/error_manager.h" +#include "utils/util.h" + +#include +#include +#include +#include + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +namespace { + +std::string ToLower(std::string_view value) { + std::string result(value); + std::transform(result.begin(), result.end(), result.begin(), [](unsigned char ch) { + return static_cast(std::tolower(ch)); + }); + return result; +} + +std::optional CanonicalKey(std::string_view key) { + const std::string lower = ToLower(key); + if (lower == "driver") return "Driver"; + if (lower == "description") return "Description"; + if (lower == "dsn") return "DSN"; + if (lower == "server" || lower == "endpoint") return "Endpoint"; + if (lower == "database") return "Database"; + if (lower == "authmode") return "AuthMode"; + if (lower == "token" || lower == "accesstoken") return "Token"; + if (lower == "user" || lower == "uid") return "User"; + if (lower == "password" || lower == "pwd") return "Password"; + if (lower == "metadatahost") return "MetadataHost"; + if (lower == "metadataport") return "MetadataPort"; + if (lower == "serviceaccountkeyfile" || lower == "safile") return "ServiceAccountKeyFile"; + if (lower == "oauth2keyfile") return "OAuth2KeyFile"; + if (lower == "iamendpoint") return "IamEndpoint"; + if (lower == "rootcertificate" || lower == "cafile") return "RootCertificate"; + if (lower == "clientcertificate") return "ClientCertificate"; + if (lower == "clientprivatekey") return "ClientPrivateKey"; + return std::nullopt; +} + +[[noreturn]] void ThrowInvalidAttribute(std::string_view attribute, std::string_view detail) { + throw TOdbcException("HY024", 0, "Invalid connection string attribute " + + std::string(attribute) + ": " + std::string(detail)); +} + +bool Has(const TConnectionParameters& parameters, std::string_view key) { + return parameters.contains(std::string(key)); +} + +std::string_view Get(const TConnectionParameters& parameters, std::string_view key) { + const auto it = parameters.find(std::string(key)); + return it == parameters.end() ? std::string_view{} : std::string_view(it->second); +} + +std::string_view RequireNonEmpty( + const TConnectionParameters& parameters, + std::string_view key, + std::string_view authMode) +{ + const auto value = Get(parameters, key); + if (value.empty()) { + throw TOdbcException("28000", 0, std::string(authMode) + + " authentication requires " + std::string(key)); + } + return value; +} + +std::string ReadDsnValue(std::string_view dsn, std::string_view key) { + const std::string dsnName(dsn); + const std::string attribute(key); + std::vector buffer(256); + while (buffer.size() <= 1024 * 1024) { + const int length = SQLGetPrivateProfileString( + dsnName.c_str(), attribute.c_str(), "", buffer.data(), static_cast(buffer.size()), nullptr); + if (length < 0) { + return {}; + } + if (static_cast(length) + 1 < buffer.size()) { + return std::string(buffer.data(), static_cast(length)); + } + buffer.resize(buffer.size() * 2); + } + throw TOdbcException("08001", 0, "DSN attribute is too large: " + attribute); +} + +std::string ReadFile(std::string_view attribute, std::string_view path) { + const std::string pathString(path); + std::ifstream input(pathString, std::ios::binary); + if (!input) { + throw TOdbcException("08001", 0, "Unable to read " + std::string(attribute) + " file: " + pathString); + } + std::string content{ + std::istreambuf_iterator(input), + std::istreambuf_iterator()}; + if (content.empty()) { + throw TOdbcException("08001", 0, std::string(attribute) + " file is empty: " + pathString); + } + return content; +} + +struct TEndpointSettings { + std::string Endpoint; + bool Secure = false; + bool ExplicitlyInsecure = false; +}; + +TEndpointSettings ParseYdbEndpoint(std::string_view value) { + constexpr std::string_view grpc = "grpc://"; + constexpr std::string_view grpcs = "grpcs://"; + if (value.starts_with(grpc)) { + return {std::string(value.substr(grpc.size())), false, true}; + } + if (value.starts_with(grpcs)) { + return {std::string(value.substr(grpcs.size())), true, false}; + } + if (value.find("://") != std::string::npos) { + ThrowInvalidAttribute("Endpoint", "only grpc:// and grpcs:// protocols are supported"); + } + return {std::string(value), false, false}; +} + +void ApplyIamEndpoint(TIamJwtFilename& params, std::string_view value) { + if (value.empty()) { + return; + } + constexpr std::string_view grpc = "grpc://"; + constexpr std::string_view grpcs = "grpcs://"; + if (value.starts_with(grpc)) { + params.Endpoint = std::string(value.substr(grpc.size())); + params.EnableSsl = false; + } else if (value.starts_with(grpcs)) { + params.Endpoint = std::string(value.substr(grpcs.size())); + params.EnableSsl = true; + } else if (value.find("://") != std::string::npos) { + ThrowInvalidAttribute("IamEndpoint", "service-account IAM supports grpc:// and grpcs://"); + } else { + params.Endpoint = std::string(value); + params.EnableSsl = true; + } +} + +uint32_t ParseMetadataPort(std::string_view value) { + if (value.empty()) { + ThrowInvalidAttribute("MetadataPort", "value is empty"); + } + uint32_t port = 0; + const auto [end, error] = std::from_chars(value.data(), value.data() + value.size(), port); + if (error != std::errc() || end != value.data() + value.size() || port == 0 || port > 65535) { + ThrowInvalidAttribute("MetadataPort", "expected an integer from 1 to 65535"); + } + return port; +} + +EAuthenticationMode ParseAuthMode(std::string_view value) { + const std::string mode = ToLower(value); + if (mode == "anonymous") return EAuthenticationMode::Anonymous; + if (mode == "token") return EAuthenticationMode::Token; + if (mode == "static") return EAuthenticationMode::Static; + if (mode == "metadata") return EAuthenticationMode::Metadata; + if (mode == "serviceaccount") return EAuthenticationMode::ServiceAccount; + if (mode == "oauth2") return EAuthenticationMode::OAuth2; + if (mode == "environment") return EAuthenticationMode::Environment; + throw TOdbcException("28000", 0, "Unknown authentication mode: " + std::string(value)); +} + +EAuthenticationMode ResolveAuthMode(const TConnectionParameters& parameters) { + const bool token = Has(parameters, "Token"); + const bool staticCredentials = Has(parameters, "User") || Has(parameters, "Password"); + const bool metadata = Has(parameters, "MetadataHost") || Has(parameters, "MetadataPort"); + const bool serviceAccount = Has(parameters, "ServiceAccountKeyFile"); + const bool oauth2 = Has(parameters, "OAuth2KeyFile"); + const size_t familyCount = static_cast(token) + static_cast(staticCredentials) + + static_cast(metadata) + static_cast(serviceAccount) + static_cast(oauth2); + + EAuthenticationMode mode; + if (Has(parameters, "AuthMode")) { + mode = ParseAuthMode(Get(parameters, "AuthMode")); + } else if (familyCount == 0) { + if (Has(parameters, "IamEndpoint")) { + throw TOdbcException("28000", 0, "IamEndpoint requires ServiceAccount or OAuth2 authentication"); + } + mode = EAuthenticationMode::Anonymous; + } else if (familyCount > 1) { + throw TOdbcException("28000", 0, "Authentication mode is ambiguous"); + } else if (token) { + mode = EAuthenticationMode::Token; + } else if (staticCredentials) { + mode = EAuthenticationMode::Static; + } else if (metadata) { + mode = EAuthenticationMode::Metadata; + } else if (serviceAccount) { + mode = EAuthenticationMode::ServiceAccount; + } else { + mode = EAuthenticationMode::OAuth2; + } + + const bool modeMatchesFamily = + (mode == EAuthenticationMode::Token && token && familyCount == 1) || + (mode == EAuthenticationMode::Static && staticCredentials && familyCount == 1) || + (mode == EAuthenticationMode::Metadata && (!familyCount || (metadata && familyCount == 1))) || + (mode == EAuthenticationMode::ServiceAccount && serviceAccount && familyCount == 1) || + (mode == EAuthenticationMode::OAuth2 && oauth2 && familyCount == 1) || + ((mode == EAuthenticationMode::Anonymous || mode == EAuthenticationMode::Environment) && familyCount == 0); + if (!modeMatchesFamily) { + throw TOdbcException("28000", 0, "Credential attributes conflict with the selected authentication mode"); + } + if (Has(parameters, "IamEndpoint") && mode != EAuthenticationMode::ServiceAccount && + mode != EAuthenticationMode::OAuth2) { + throw TOdbcException("28000", 0, "IamEndpoint is valid only for ServiceAccount or OAuth2 authentication"); + } + return mode; +} + +} // namespace + +TConnectionParameters ParseAndNormalizeConnectionString( + std::string_view connectionString, + std::vector& ignoredAttributes) +{ + TConnectionParameters parameters; + for (const auto& [key, value] : ParseConnectionStringEntries(connectionString)) { + const auto canonical = CanonicalKey(key); + if (!canonical) { + ignoredAttributes.push_back(key); + continue; + } + parameters[*canonical] = value; + } + return parameters; +} + +TConnectionParameters ReadDsnParameters(std::string_view dsn) { + TConnectionParameters parameters; + // Aliases are read first so the canonical spelling wins inside a DSN. + static constexpr std::array keys = { + "Server", "UID", "PWD", "AccessToken", "SaFile", "CaFile", + "Driver", "Description", "Endpoint", "Database", "AuthMode", "Token", + "User", "Password", "MetadataHost", "MetadataPort", "ServiceAccountKeyFile", + "OAuth2KeyFile", "IamEndpoint", "RootCertificate", "ClientCertificate", + "ClientPrivateKey", "DSN"}; + for (const char* key : keys) { + std::string value = ReadDsnValue(dsn, key); + if (!value.empty()) { + parameters[*CanonicalKey(key)] = std::move(value); + } + } + return parameters; +} + +void OverlayConnectionParameters(TConnectionParameters& destination, const TConnectionParameters& source) { + std::optional selectedMode; + if (Has(source, "AuthMode")) { + selectedMode = ParseAuthMode(Get(source, "AuthMode")); + } else { + const bool token = Has(source, "Token"); + const bool staticCredentials = Has(source, "User") || Has(source, "Password"); + const bool metadata = Has(source, "MetadataHost") || Has(source, "MetadataPort"); + const bool serviceAccount = Has(source, "ServiceAccountKeyFile"); + const bool oauth2 = Has(source, "OAuth2KeyFile"); + const size_t familyCount = static_cast(token) + static_cast(staticCredentials) + + static_cast(metadata) + static_cast(serviceAccount) + static_cast(oauth2); + if (familyCount == 1) { + selectedMode = token ? EAuthenticationMode::Token + : staticCredentials ? EAuthenticationMode::Static + : metadata ? EAuthenticationMode::Metadata + : serviceAccount ? EAuthenticationMode::ServiceAccount + : EAuthenticationMode::OAuth2; + destination.erase("AuthMode"); + } + } + + if (selectedMode) { + const auto belongsToSelectedMode = [selectedMode](std::string_view key) { + switch (*selectedMode) { + case EAuthenticationMode::Token: + return key == "Token"; + case EAuthenticationMode::Static: + return key == "User" || key == "Password"; + case EAuthenticationMode::Metadata: + return key == "MetadataHost" || key == "MetadataPort"; + case EAuthenticationMode::ServiceAccount: + return key == "ServiceAccountKeyFile" || key == "IamEndpoint"; + case EAuthenticationMode::OAuth2: + return key == "OAuth2KeyFile" || key == "IamEndpoint"; + case EAuthenticationMode::Anonymous: + case EAuthenticationMode::Environment: + return false; + } + return false; + }; + for (const std::string_view key : { + "Token", "User", "Password", "MetadataHost", "MetadataPort", + "ServiceAccountKeyFile", "OAuth2KeyFile", "IamEndpoint"}) { + if (!belongsToSelectedMode(key)) { + destination.erase(std::string(key)); + } + } + } + + for (const auto& [key, value] : source) { + destination[key] = value; + } +} + +TResolvedConnectionSettings ResolveConnectionSettings( + TConnectionParameters parameters, + std::string dataSourceName) +{ + const std::string endpointValue(Get(parameters, "Endpoint")); + const std::string database(Get(parameters, "Database")); + if (endpointValue.empty() || database.empty()) { + throw TOdbcException("08001", 0, "Missing Endpoint (or Server) or Database"); + } + + const TEndpointSettings endpoint = ParseYdbEndpoint(endpointValue); + const bool hasRoot = Has(parameters, "RootCertificate"); + const bool hasClientCert = Has(parameters, "ClientCertificate"); + const bool hasClientKey = Has(parameters, "ClientPrivateKey"); + if (hasClientCert != hasClientKey) { + throw TOdbcException("08001", 0, + "ClientCertificate and ClientPrivateKey must be specified together"); + } + const bool hasTlsFiles = hasRoot || hasClientCert; + if (endpoint.ExplicitlyInsecure && hasTlsFiles) { + ThrowInvalidAttribute("Endpoint", "grpc:// cannot be combined with TLS certificate attributes"); + } + + const EAuthenticationMode authMode = ResolveAuthMode(parameters); + TDriverConfig driverConfig = authMode == EAuthenticationMode::Environment + ? CreateFromEnvironment() + : TDriverConfig(); + driverConfig.SetEndpoint(endpoint.Endpoint).SetDatabase(database); + + switch (authMode) { + case EAuthenticationMode::Anonymous: + driverConfig.SetCredentialsProviderFactory(CreateInsecureCredentialsProviderFactory()); + break; + case EAuthenticationMode::Token: + driverConfig.SetCredentialsProviderFactory(CreateOAuthCredentialsProviderFactory( + std::string(RequireNonEmpty(parameters, "Token", "Token")))); + break; + case EAuthenticationMode::Static: + driverConfig.SetCredentialsProviderFactory(CreateLoginCredentialsProviderFactory({ + .User = std::string(RequireNonEmpty(parameters, "User", "Static")), + .Password = std::string(RequireNonEmpty(parameters, "Password", "Static")), + })); + break; + case EAuthenticationMode::Metadata: { + TIamHost params; + if (Has(parameters, "MetadataHost")) { + params.Host = std::string(RequireNonEmpty(parameters, "MetadataHost", "Metadata")); + } + if (Has(parameters, "MetadataPort")) { + params.Port = ParseMetadataPort(Get(parameters, "MetadataPort")); + } + driverConfig.SetCredentialsProviderFactory(CreateIamCredentialsProviderFactory(params)); + break; + } + case EAuthenticationMode::ServiceAccount: { + TIamJwtFilename params; + params.JwtFilename = std::string(RequireNonEmpty(parameters, "ServiceAccountKeyFile", "ServiceAccount")); + ApplyIamEndpoint(params, Get(parameters, "IamEndpoint")); + try { + driverConfig.SetCredentialsProviderFactory(CreateIamJwtFileCredentialsProviderFactory(params)); + } catch (const std::exception& ex) { + throw TOdbcException("08001", 0, + "Unable to load ServiceAccountKeyFile " + params.JwtFilename + ": " + ex.what()); + } + break; + } + case EAuthenticationMode::OAuth2: { + const std::string path(RequireNonEmpty(parameters, "OAuth2KeyFile", "OAuth2")); + try { + driverConfig.SetCredentialsProviderFactory( + CreateOauth2TokenExchangeFileCredentialsProviderFactory( + path, std::string(Get(parameters, "IamEndpoint")))); + } catch (const std::exception& ex) { + throw TOdbcException("08001", 0, + "Unable to load OAuth2KeyFile " + path + ": " + ex.what()); + } + break; + } + case EAuthenticationMode::Environment: + break; + } + + const bool secure = endpoint.Secure || hasTlsFiles; + std::string rootPem; + std::string clientCertPem; + std::string clientKeyPem; + if (hasRoot) { + rootPem = ReadFile("RootCertificate", Get(parameters, "RootCertificate")); + } + if (hasClientCert) { + clientCertPem = ReadFile("ClientCertificate", Get(parameters, "ClientCertificate")); + clientKeyPem = ReadFile("ClientPrivateKey", Get(parameters, "ClientPrivateKey")); + } + if (secure) { + driverConfig.UseSecureConnection(rootPem); + } + if (hasClientCert) { + driverConfig.UseClientCertificate(clientCertPem, clientKeyPem); + } + + if (dataSourceName.empty()) { + dataSourceName = std::string(Get(parameters, "DSN")); + } + return { + .Endpoint = endpoint.Endpoint, + .Database = database, + .DataSourceName = std::move(dataSourceName), + .DriverConfig = std::move(driverConfig), + }; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/connection_config.h b/odbc/src/connection_config.h new file mode 100644 index 00000000000..2caa640e504 --- /dev/null +++ b/odbc/src/connection_config.h @@ -0,0 +1,41 @@ +#pragma once + +#include + +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +enum class EAuthenticationMode { + Anonymous, + Token, + Static, + Metadata, + ServiceAccount, + OAuth2, + Environment, +}; + +using TConnectionParameters = std::map; + +struct TResolvedConnectionSettings { + std::string Endpoint; + std::string Database; + std::string DataSourceName; + TDriverConfig DriverConfig; +}; + +TConnectionParameters ParseAndNormalizeConnectionString( + std::string_view connectionString, + std::vector& ignoredAttributes); +TConnectionParameters ReadDsnParameters(std::string_view dsn); +void OverlayConnectionParameters(TConnectionParameters& destination, const TConnectionParameters& source); + +TResolvedConnectionSettings ResolveConnectionSettings( + TConnectionParameters parameters, + std::string dataSourceName = {}); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/descriptor.cpp b/odbc/src/descriptor.cpp new file mode 100644 index 00000000000..111fb5ff0b2 --- /dev/null +++ b/odbc/src/descriptor.cpp @@ -0,0 +1,319 @@ +#include "descriptor.h" + +#include "statement.h" +#include "utils/diag.h" +#include "utils/sql_type_map.h" + +#include +#include + +namespace NYdb::NOdbc { +namespace { + +bool IsCharacter(SQLSMALLINT type) { + return type == SQL_CHAR || type == SQL_VARCHAR || type == SQL_LONGVARCHAR + || type == SQL_WCHAR || type == SQL_WVARCHAR || type == SQL_WLONGVARCHAR; +} + +SQLSMALLINT DateTimeCode(SQLSMALLINT type) { + switch (type) { + case SQL_TYPE_DATE: return SQL_CODE_DATE; + case SQL_TYPE_TIME: return SQL_CODE_TIME; + case SQL_TYPE_TIMESTAMP: return SQL_CODE_TIMESTAMP; + default: return 0; + } +} + +std::string TypeName(SQLSMALLINT type) { + const TSqlTypeSpec* spec = FindSqlTypeSpec(type); + return spec ? std::string(spec->Name) : type == SQL_GUID ? "GUID" : ""; +} + +SQLRETURN WriteString(TErrorManager& errors, const std::string& text, SQLPOINTER value, + SQLINTEGER bufferLength, SQLINTEGER* lengthPtr) { + if (lengthPtr) { + *lengthPtr = static_cast(text.size()); + } + if (!value) { + return SQL_SUCCESS; + } + if (bufferLength < 0) { + return Diag::AddInvalidBufferLength(errors); + } + if (bufferLength == 0) { + return text.empty() ? SQL_SUCCESS : Diag::AddRightTruncated(errors); + } + const auto copyLength = std::min(text.size(), static_cast(bufferLength - 1)); + std::memcpy(value, text.data(), copyLength); + static_cast(value)[copyLength] = '\0'; + return copyLength == text.size() ? SQL_SUCCESS : Diag::AddRightTruncated(errors); +} + +template +SQLRETURN WriteScalar(SQLPOINTER value, T scalar) { + *static_cast(value) = scalar; + return SQL_SUCCESS; +} + +} // namespace + +TDescriptor::TDescriptor(EDescType type, TConnection* conn) + : Type_(type) + , Conn_(conn) { + if (Type_ == EDescType::Explicit) { + Conn_->RegisterDescriptor(this); + } +} + +TDescriptor::~TDescriptor() { + while (!Statements_.empty()) { + Statements_.back()->DetachDescriptor(this); + } + if (Type_ == EDescType::Explicit) { + Conn_->UnregisterDescriptor(this); + } +} + +TDescriptor* TDescriptor::FromHandle(SQLHDESC handle) { + if (!handle) { + throw TOdbcException("HY000", 0, "Invalid handle", SQL_INVALID_HANDLE); + } + return static_cast(handle); +} + +TDescRecord& TDescriptor::Record(SQLSMALLINT number) { + if (number < 1) { + throw TOdbcException("07009", 0, "Invalid descriptor index"); + } + if (Records_.size() < static_cast(number)) { + Records_.resize(static_cast(number)); + } + TDescRecord& record = Records_[static_cast(number - 1)]; + record.Active = true; + return record; +} + +const TDescRecord* TDescriptor::FindRecord(SQLSMALLINT number) const noexcept { + if (number < 1 || static_cast(number) > Records_.size()) { + return nullptr; + } + const TDescRecord& record = Records_[static_cast(number - 1)]; + return record.Active ? &record : nullptr; +} + +TDescRecord* TDescriptor::FindRecord(SQLSMALLINT number) noexcept { + return const_cast(std::as_const(*this).FindRecord(number)); +} + +void TDescriptor::RemoveRecord(SQLSMALLINT number) { + if (number > 0 && static_cast(number) <= Records_.size()) { + Records_[static_cast(number - 1)] = {}; + while (!Records_.empty() && !Records_.back().Active) { + Records_.pop_back(); + } + } +} + +SQLSMALLINT TDescriptor::GetRecordCount() const noexcept { + return static_cast(Records_.size()); +} + +void TDescriptor::Attach(TStatement* stmt) { + if (Type_ == EDescType::Explicit + && std::find(Statements_.begin(), Statements_.end(), stmt) == Statements_.end()) { + Statements_.push_back(stmt); + } +} + +void TDescriptor::Detach(TStatement* stmt) { + std::erase(Statements_, stmt); +} + +SQLRETURN TDescriptor::GetDescField(SQLSMALLINT recNumber, SQLSMALLINT field, SQLPOINTER value, + SQLINTEGER bufferLength, SQLINTEGER* lengthPtr) { + const bool stringField = field == SQL_DESC_BASE_COLUMN_NAME || field == SQL_DESC_NAME + || field == SQL_DESC_TYPE_NAME || field == SQL_DESC_LOCAL_TYPE_NAME + || field == SQL_DESC_LITERAL_PREFIX || field == SQL_DESC_LITERAL_SUFFIX; + if (!value && (!stringField || !lengthPtr)) { + return Diag::AddNullPointer(*this); + } + switch (field) { + case SQL_DESC_ALLOC_TYPE: + return WriteScalar(value, static_cast( + Type_ == EDescType::Explicit ? SQL_DESC_ALLOC_USER : SQL_DESC_ALLOC_AUTO)); + case SQL_DESC_COUNT: return WriteScalar(value, GetRecordCount()); + case SQL_DESC_ARRAY_SIZE: return WriteScalar(value, ArraySize_); + case SQL_DESC_BIND_TYPE: return WriteScalar(value, BindType_); + case SQL_DESC_BIND_OFFSET_PTR: return WriteScalar(value, BindOffsetPtr_); + case SQL_DESC_ARRAY_STATUS_PTR: return WriteScalar(value, ArrayStatusPtr_); + case SQL_DESC_ROWS_PROCESSED_PTR: return WriteScalar(value, RowsProcessedPtr_); + default: break; + } + + const TDescRecord* record = FindRecord(recNumber); + if (!record) { + return recNumber > GetRecordCount() + ? SQL_NO_DATA + : AddError("07009", 0, "Invalid descriptor index"); + } + switch (field) { + case SQL_DESC_BASE_COLUMN_NAME: + case SQL_DESC_NAME: + return WriteString(*this, record->Name, value, bufferLength, lengthPtr); + case SQL_DESC_TYPE_NAME: + case SQL_DESC_LOCAL_TYPE_NAME: + return WriteString(*this, TypeName(record->Type), value, bufferLength, lengthPtr); + case SQL_DESC_LITERAL_PREFIX: + case SQL_DESC_LITERAL_SUFFIX: + return WriteString( + *this, IsCharacter(record->Type) || DateTimeCode(record->Type) ? "'" : "", + value, bufferLength, lengthPtr); + case SQL_DESC_TYPE: + return WriteScalar(value, static_cast( + DateTimeCode(record->Type) ? SQL_DATETIME : record->Type)); + case SQL_DESC_CONCISE_TYPE: return WriteScalar(value, record->Type); + case SQL_DESC_DATETIME_INTERVAL_CODE: + return WriteScalar(value, DateTimeCode(record->Type)); + case SQL_DESC_LENGTH: return WriteScalar(value, record->Length); + case SQL_DESC_OCTET_LENGTH: return WriteScalar(value, record->OctetLength); + case SQL_DESC_DISPLAY_SIZE: return WriteScalar(value, record->Length); + case SQL_DESC_PRECISION: return WriteScalar(value, record->Precision); + case SQL_DESC_SCALE: return WriteScalar(value, record->Scale); + case SQL_DESC_NULLABLE: return WriteScalar(value, record->Nullable); + case SQL_DESC_CASE_SENSITIVE: + return WriteScalar(value, static_cast(IsCharacter(record->Type))); + case SQL_DESC_FIXED_PREC_SCALE: + return WriteScalar(value, static_cast( + record->Type == SQL_DECIMAL || record->Type == SQL_NUMERIC)); + case SQL_DESC_SEARCHABLE: + return WriteScalar(value, static_cast(SQL_SEARCHABLE)); + case SQL_DESC_UNNAMED: + return WriteScalar(value, static_cast( + record->Name.empty() ? SQL_UNNAMED : SQL_NAMED)); + case SQL_DESC_UNSIGNED: return WriteScalar(value, static_cast(SQL_FALSE)); + case SQL_DESC_UPDATABLE: + return WriteScalar(value, static_cast(SQL_ATTR_READONLY)); + case SQL_DESC_PARAMETER_TYPE: return WriteScalar(value, record->ParameterType); + case SQL_DESC_DATA_PTR: return WriteScalar(value, record->DataPtr); + case SQL_DESC_INDICATOR_PTR: return WriteScalar(value, record->IndicatorPtr); + case SQL_DESC_OCTET_LENGTH_PTR: return WriteScalar(value, record->OctetLengthPtr); + default: return Diag::AddNotImplemented(*this); + } +} + +SQLRETURN TDescriptor::GetDescRec(SQLSMALLINT recNumber, SQLCHAR* name, SQLSMALLINT bufferLength, + SQLSMALLINT* nameLengthPtr, SQLSMALLINT* typePtr, + SQLSMALLINT* subTypePtr, SQLLEN* lengthPtr, + SQLSMALLINT* precisionPtr, SQLSMALLINT* scalePtr, + SQLSMALLINT* nullablePtr) { + const TDescRecord* record = FindRecord(recNumber); + if (!record) { + return recNumber > GetRecordCount() + ? SQL_NO_DATA + : AddError("07009", 0, "Invalid descriptor index"); + } + SQLINTEGER nameLength = 0; + SQLRETURN result = SQL_SUCCESS; + if (name) { + result = WriteString(*this, record->Name, name, bufferLength, &nameLength); + } else { + nameLength = static_cast(record->Name.size()); + } + if (nameLengthPtr) *nameLengthPtr = static_cast(nameLength); + if (typePtr) *typePtr = DateTimeCode(record->Type) ? SQL_DATETIME : record->Type; + if (subTypePtr) *subTypePtr = DateTimeCode(record->Type); + if (lengthPtr) *lengthPtr = record->Length; + if (precisionPtr) *precisionPtr = record->Precision; + if (scalePtr) *scalePtr = record->Scale; + if (nullablePtr) *nullablePtr = record->Nullable; + return result; +} + +SQLRETURN TDescriptor::SetDescField(SQLSMALLINT recNumber, SQLSMALLINT field, SQLPOINTER value, + SQLINTEGER bufferLength) { + switch (field) { + case SQL_DESC_COUNT: { + const auto count = static_cast(reinterpret_cast(value)); + if (count < 0) return AddError("HY024", 0, "Invalid SQL_DESC_COUNT value"); + Records_.resize(static_cast(count)); + for (auto& record : Records_) record.Active = true; + return SQL_SUCCESS; + } + case SQL_DESC_ARRAY_SIZE: { + const auto size = static_cast(reinterpret_cast(value)); + if (size == 0) return AddError("HY024", 0, "Invalid SQL_DESC_ARRAY_SIZE value"); + ArraySize_ = size; + return SQL_SUCCESS; + } + case SQL_DESC_BIND_TYPE: + BindType_ = static_cast(reinterpret_cast(value)); + return SQL_SUCCESS; + case SQL_DESC_BIND_OFFSET_PTR: BindOffsetPtr_ = static_cast(value); return SQL_SUCCESS; + case SQL_DESC_ARRAY_STATUS_PTR: ArrayStatusPtr_ = static_cast(value); return SQL_SUCCESS; + case SQL_DESC_ROWS_PROCESSED_PTR: RowsProcessedPtr_ = static_cast(value); return SQL_SUCCESS; + default: break; + } + + if (Type_ == EDescType::ImpRow) { + return AddError("HY016", 0, "Cannot modify an implementation row descriptor"); + } + + TDescRecord& record = Record(recNumber); + const auto integer = reinterpret_cast(value); + switch (field) { + case SQL_DESC_TYPE: + case SQL_DESC_CONCISE_TYPE: record.Type = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_LENGTH: record.Length = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_OCTET_LENGTH: record.OctetLength = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_PRECISION: record.Precision = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_SCALE: record.Scale = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_NULLABLE: record.Nullable = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_PARAMETER_TYPE: record.ParameterType = static_cast(integer); return SQL_SUCCESS; + case SQL_DESC_DATA_PTR: record.DataPtr = value; return SQL_SUCCESS; + case SQL_DESC_INDICATOR_PTR: record.IndicatorPtr = static_cast(value); return SQL_SUCCESS; + case SQL_DESC_OCTET_LENGTH_PTR: record.OctetLengthPtr = static_cast(value); return SQL_SUCCESS; + case SQL_DESC_NAME: + if (!value) return Diag::AddNullPointer(*this); + record.Name = bufferLength == SQL_NTS + ? std::string(static_cast(value)) + : std::string(static_cast(value), static_cast(bufferLength)); + return SQL_SUCCESS; + default: return Diag::AddNotImplemented(*this); + } +} + +SQLRETURN TDescriptor::SetDescRec(SQLSMALLINT recNumber, SQLSMALLINT type, SQLSMALLINT subType, + SQLLEN length, SQLSMALLINT precision, SQLSMALLINT scale, + SQLPOINTER dataPtr, SQLLEN* stringLengthPtr, + SQLLEN* indicatorPtr) { + if (Type_ == EDescType::ImpRow) { + return AddError("HY016", 0, "Cannot modify an implementation row descriptor"); + } + TDescRecord& record = Record(recNumber); + record.Type = type; + record.SubType = subType; + record.Length = length; + record.OctetLength = length; + record.Precision = precision; + record.Scale = scale; + record.DataPtr = dataPtr; + record.OctetLengthPtr = stringLengthPtr; + record.IndicatorPtr = indicatorPtr; + return SQL_SUCCESS; +} + +SQLRETURN TDescriptor::CopyDesc(TDescriptor* target) { + if (!target) return Diag::AddNullPointer(*this); + if (target->Type_ == EDescType::ImpRow) { + return AddError("HY016", 0, "Cannot modify an implementation row descriptor"); + } + target->ArraySize_ = ArraySize_; + target->BindType_ = BindType_; + target->BindOffsetPtr_ = BindOffsetPtr_; + target->ArrayStatusPtr_ = ArrayStatusPtr_; + target->RowsProcessedPtr_ = RowsProcessedPtr_; + target->Records_ = Records_; + return SQL_SUCCESS; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/descriptor.h b/odbc/src/descriptor.h new file mode 100644 index 00000000000..3125ac5b80d --- /dev/null +++ b/odbc/src/descriptor.h @@ -0,0 +1,100 @@ +#pragma once + +#include "utils/error_manager.h" + +#include +#include + +#include +#include + +namespace NYdb::NOdbc { + +class TConnection; +class TStatement; + +enum class EDescType { + AppRow, + AppParam, + ImpRow, + ImpParam, + Explicit, +}; + +struct TDescRecord { + std::string Name; + SQLSMALLINT Type = SQL_UNKNOWN_TYPE; + SQLSMALLINT SubType = 0; + SQLLEN Length = 0; + SQLLEN OctetLength = 0; + SQLSMALLINT Precision = 0; + SQLSMALLINT Scale = 0; + SQLSMALLINT Nullable = SQL_NULLABLE; + SQLPOINTER DataPtr = nullptr; + SQLLEN* IndicatorPtr = nullptr; + SQLLEN* OctetLengthPtr = nullptr; + SQLSMALLINT ParameterType = SQL_PARAM_INPUT; + bool Active = false; + bool AtExec = false; + bool AtExecComplete = false; + SQLLEN AtExecIndicator = 0; + std::string AtExecChunk; +}; + +class TDescriptor : public TErrorManager { +public: + TDescriptor(EDescType type, TConnection* conn); + ~TDescriptor(); + + EDescType GetDescType() const noexcept { return Type_; } + TConnection* GetConnection() const noexcept { return Conn_; } + + SQLULEN GetArraySize() const noexcept { return ArraySize_; } + SQLULEN GetBindType() const noexcept { return BindType_; } + SQLULEN* GetBindOffsetPtr() const noexcept { return BindOffsetPtr_; } + SQLUSMALLINT* GetArrayStatusPtr() const noexcept { return ArrayStatusPtr_; } + SQLULEN* GetRowsProcessedPtr() const noexcept { return RowsProcessedPtr_; } + void SetArraySize(SQLULEN value) noexcept { ArraySize_ = value; } + void SetBindType(SQLULEN value) noexcept { BindType_ = value; } + void SetBindOffsetPtr(SQLULEN* value) noexcept { BindOffsetPtr_ = value; } + void SetArrayStatusPtr(SQLUSMALLINT* value) noexcept { ArrayStatusPtr_ = value; } + void SetRowsProcessedPtr(SQLULEN* value) noexcept { RowsProcessedPtr_ = value; } + + TDescRecord& Record(SQLSMALLINT number); + const TDescRecord* FindRecord(SQLSMALLINT number) const noexcept; + TDescRecord* FindRecord(SQLSMALLINT number) noexcept; + void RemoveRecord(SQLSMALLINT number); + void ClearRecords() noexcept { Records_.clear(); } + SQLSMALLINT GetRecordCount() const noexcept; + + void Attach(TStatement* stmt); + void Detach(TStatement* stmt); + + SQLRETURN GetDescField(SQLSMALLINT recNumber, SQLSMALLINT fieldIdentifier, SQLPOINTER value, + SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr); + SQLRETURN GetDescRec(SQLSMALLINT recNumber, SQLCHAR* name, SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr, SQLSMALLINT* typePtr, SQLSMALLINT* subTypePtr, + SQLLEN* lengthPtr, SQLSMALLINT* precisionPtr, SQLSMALLINT* scalePtr, + SQLSMALLINT* nullablePtr); + SQLRETURN SetDescField(SQLSMALLINT recNumber, SQLSMALLINT fieldIdentifier, SQLPOINTER value, + SQLINTEGER bufferLength); + SQLRETURN SetDescRec(SQLSMALLINT recNumber, SQLSMALLINT type, SQLSMALLINT subType, SQLLEN length, + SQLSMALLINT precision, SQLSMALLINT scale, SQLPOINTER dataPtr, + SQLLEN* stringLengthPtr, SQLLEN* indicatorPtr); + SQLRETURN CopyDesc(TDescriptor* target); + + static TDescriptor* FromHandle(SQLHDESC handle); + +private: + EDescType Type_; + TConnection* Conn_; + SQLULEN ArraySize_ = 1; + SQLULEN BindType_ = SQL_BIND_BY_COLUMN; + SQLULEN* BindOffsetPtr_ = nullptr; + SQLUSMALLINT* ArrayStatusPtr_ = nullptr; + SQLULEN* RowsProcessedPtr_ = nullptr; + std::vector Records_; + std::vector Statements_; +}; + +} // namespace NYdb::NOdbc diff --git a/odbc/src/environment.cpp b/odbc/src/environment.cpp new file mode 100644 index 00000000000..047ebd1fb7c --- /dev/null +++ b/odbc/src/environment.cpp @@ -0,0 +1,117 @@ +#include "environment.h" +#include "connection.h" + + #include + #include + +namespace NYdb { +namespace NOdbc { + +TEnvironment::TEnvironment() : OdbcVersion_(SQL_OV_ODBC3) {} +TEnvironment::~TEnvironment() {} + +SQLRETURN TEnvironment::SetAttribute(SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER stringLength) { + switch (attribute) { + case SQL_ATTR_ODBC_VERSION: { + if (!value) { + return AddError("HY009", 0, "Invalid use of null pointer"); + } + OdbcVersion_ = static_cast(reinterpret_cast(value)); + return SQL_SUCCESS; + } + case SQL_ATTR_OUTPUT_NTS: { + if (value && static_cast(reinterpret_cast(value)) != SQL_TRUE) { + return AddError("HY024", 0, "SQL_ATTR_OUTPUT_NTS must be SQL_TRUE"); + } + return SQL_SUCCESS; + } + default: + return AddError("HYC00", 0, "Optional feature not implemented"); + } +} + +SQLRETURN TEnvironment::GetAttribute(SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr) { + if (!value) { + return AddError("HY009", 0, "Invalid use of null pointer"); + } + if (stringLengthPtr) { + *stringLengthPtr = 0; + } + switch (attribute) { + case SQL_ATTR_ODBC_VERSION: + if (bufferLength < static_cast(sizeof(SQLINTEGER))) { + return AddError("HY090", 0, "Invalid string or buffer length"); + } + *reinterpret_cast(value) = OdbcVersion_; + if (stringLengthPtr) { + *stringLengthPtr = sizeof(SQLINTEGER); + } + return SQL_SUCCESS; + case SQL_ATTR_OUTPUT_NTS: + if (bufferLength < static_cast(sizeof(SQLINTEGER))) { + return AddError("HY090", 0, "Invalid string or buffer length"); + } + *reinterpret_cast(value) = SQL_TRUE; + if (stringLengthPtr) { + *stringLengthPtr = sizeof(SQLINTEGER); + } + return SQL_SUCCESS; + default: + return AddError("HYC00", 0, "Optional feature not implemented"); + } +} + +void TEnvironment::RegisterConnection(TConnection* conn){ + if (conn == nullptr){ + throw std::invalid_argument("null connection"); + } + Connections_.insert(conn); +} + +void TEnvironment::UnregisterConnection(TConnection* conn){ + if (conn == nullptr){ + throw std::invalid_argument("null connection"); + } + Connections_.erase(conn); +} + +std::vector TEnvironment::GetConnectionsSnapshot() const { + return std::vector(Connections_.begin(), Connections_.end()); +} + +SQLRETURN TEnvironment::EndTran(SQLSMALLINT completionType){ + if (completionType != SQL_COMMIT && completionType != SQL_ROLLBACK){ + return AddError("HY012", 0, "Invalid transaction operation code"); + } + bool hasFailures = false; + int failedCount = 0; + + for (auto* conn : Connections_) { + if (!conn || !conn->GetTx()) { + continue; + } + try { + if (completionType == SQL_COMMIT) { + conn->CommitTx(); + } else { + conn->RollbackTx(); + } + } catch (const std::exception& ex) { + hasFailures = true; + ++failedCount; + AddError("HY000", 0, ex.what()); + } catch (...) { + hasFailures = true; + ++failedCount; + AddError("HY000", 0, "Unknown error during ENV-level transaction completion"); + } + } + if (hasFailures) { + AddError("HY000", 0, + "SQLEndTran(SQL_HANDLE_ENV): " + std::to_string(failedCount) + " connection(s) failed"); + return SQL_ERROR; + } + return SQL_SUCCESS; +} +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/environment.h b/odbc/src/environment.h new file mode 100644 index 00000000000..5dc6021ce36 --- /dev/null +++ b/odbc/src/environment.h @@ -0,0 +1,35 @@ +#pragma once + +#include "utils/error_manager.h" + +#include +#include +#include +#include + +namespace NYdb { +namespace NOdbc { + +class TConnection; + +class TEnvironment : public TErrorManager { +private: + SQLINTEGER OdbcVersion_; + std::unordered_set Connections_; + +public: + TEnvironment(); + ~TEnvironment(); + + SQLRETURN SetAttribute(SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER stringLength); + SQLRETURN GetAttribute(SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr); + + void RegisterConnection(TConnection*); + void UnregisterConnection(TConnection*); + std::vector GetConnectionsSnapshot() const; + + SQLRETURN EndTran(SQLSMALLINT completionType); +}; + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/metadata.cpp b/odbc/src/metadata.cpp new file mode 100644 index 00000000000..67f2be4ec61 --- /dev/null +++ b/odbc/src/metadata.cpp @@ -0,0 +1,394 @@ +#include "metadata.h" + +#include "utils/diag.h" + +#include +#include + +namespace NYdb::NOdbc { +namespace { + +SQLRETURN WriteInfoString( + TConnection* conn, + const char* value, + SQLPOINTER infoValuePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) { + return Diag::WriteOdbcString(*conn, value, infoValuePtr, bufferLength, stringLengthPtr); +} + +template +SQLRETURN WriteInfoScalar( + TConnection* conn, + T value, + SQLPOINTER infoValuePtr, + SQLSMALLINT* stringLengthPtr) { + if (!infoValuePtr) { + return conn->AddError("HY009", 0, "Invalid use of null pointer"); + } + *reinterpret_cast(infoValuePtr) = value; + if (stringLengthPtr) { + *stringLengthPtr = static_cast(sizeof(T)); + } + return SQL_SUCCESS; +} + + +bool IsSupportedFunction(SQLUSMALLINT functionId) { + switch (functionId) { + case SQL_API_SQLALLOCHANDLE: + case SQL_API_SQLBINDCOL: + case SQL_API_SQLBINDPARAMETER: + case SQL_API_SQLCANCEL: + case SQL_API_SQLCLOSECURSOR: + case SQL_API_SQLCOLATTRIBUTE: + case SQL_API_SQLCOLUMNS: + case SQL_API_SQLCONNECT: + case SQL_API_SQLCOPYDESC: + case SQL_API_SQLDESCRIBECOL: + case SQL_API_SQLDESCRIBEPARAM: + case SQL_API_SQLDISCONNECT: + case SQL_API_SQLDRIVERCONNECT: + case SQL_API_SQLENDTRAN: + case SQL_API_SQLEXECDIRECT: + case SQL_API_SQLEXECUTE: + case SQL_API_SQLFETCH: + case SQL_API_SQLFETCHSCROLL: + case SQL_API_SQLFOREIGNKEYS: + case SQL_API_SQLFREEHANDLE: + case SQL_API_SQLFREESTMT: + case SQL_API_SQLGETCURSORNAME: + case SQL_API_SQLGETDATA: + case SQL_API_SQLGETDESCFIELD: + case SQL_API_SQLGETDESCREC: + case SQL_API_SQLGETDIAGFIELD: + case SQL_API_SQLGETDIAGREC: + case SQL_API_SQLGETFUNCTIONS: + case SQL_API_SQLGETCONNECTATTR: + case SQL_API_SQLGETENVATTR: + case SQL_API_SQLGETINFO: + case SQL_API_SQLGETSTMTATTR: + case SQL_API_SQLGETTYPEINFO: + case SQL_API_SQLMORERESULTS: + case SQL_API_SQLNATIVESQL: + case SQL_API_SQLNUMPARAMS: + case SQL_API_SQLNUMRESULTCOLS: + case SQL_API_SQLPARAMDATA: + case SQL_API_SQLPREPARE: + case SQL_API_SQLPRIMARYKEYS: + case SQL_API_SQLPUTDATA: + case SQL_API_SQLROWCOUNT: + case SQL_API_SQLSETCONNECTATTR: + case SQL_API_SQLSETCURSORNAME: + case SQL_API_SQLSETDESCFIELD: + case SQL_API_SQLSETDESCREC: + case SQL_API_SQLSETENVATTR: + case SQL_API_SQLSETSTMTATTR: + case SQL_API_SQLSPECIALCOLUMNS: + case SQL_API_SQLSTATISTICS: + case SQL_API_SQLTABLES: + return true; + default: + return false; + } +} + +} // namespace + +SQLRETURN NMetadata::GetInfo( + TConnection* conn, + SQLUSMALLINT infoType, + SQLPOINTER infoValuePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) { + switch (infoType) { + // Driver Information + case SQL_DRIVER_NAME: + return WriteInfoString(conn, "ydb-odbc", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_DRIVER_VER: + return WriteInfoString(conn, "unknown", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_DRIVER_ODBC_VER: + return WriteInfoString(conn, "03.00", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_ODBC_INTERFACE_CONFORMANCE: + return WriteInfoScalar(conn, SQL_OIC_CORE, infoValuePtr, stringLengthPtr); + case SQL_ODBC_API_CONFORMANCE: + return WriteInfoScalar(conn, SQL_OAC_LEVEL1, infoValuePtr, stringLengthPtr); + case SQL_ODBC_SAG_CLI_CONFORMANCE: + return WriteInfoScalar(conn, SQL_OSCC_NOT_COMPLIANT, infoValuePtr, stringLengthPtr); + case SQL_ODBC_SQL_CONFORMANCE: + return WriteInfoScalar(conn, SQL_OSC_MINIMUM, infoValuePtr, stringLengthPtr); + case SQL_MAX_TABLE_NAME_LEN: + case SQL_MAX_COLUMN_NAME_LEN: + case SQL_MAX_CATALOG_NAME_LEN: + case SQL_MAX_IDENTIFIER_LEN: + return WriteInfoScalar(conn, 255, infoValuePtr, stringLengthPtr); + case SQL_MAX_SCHEMA_NAME_LEN: + case SQL_MAX_PROCEDURE_NAME_LEN: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + case SQL_MAX_USER_NAME_LEN: + return WriteInfoScalar(conn, 128, infoValuePtr, stringLengthPtr); + case SQL_MAX_DRIVER_CONNECTIONS: + case SQL_MAX_CONCURRENT_ACTIVITIES: + case SQL_MAX_STATEMENT_LEN: + case SQL_MAX_BINARY_LITERAL_LEN: + case SQL_MAX_CHAR_LITERAL_LEN: + case SQL_MAX_COLUMNS_IN_GROUP_BY: + case SQL_MAX_COLUMNS_IN_ORDER_BY: + case SQL_MAX_COLUMNS_IN_INDEX: + case SQL_MAX_COLUMNS_IN_SELECT: + case SQL_MAX_COLUMNS_IN_TABLE: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + case SQL_SEARCH_PATTERN_ESCAPE: + return WriteInfoString(conn, "\\", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_KEYWORDS: + case SQL_SPECIAL_CHARACTERS: + return WriteInfoString(conn, "", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_CONCAT_NULL_BEHAVIOR: + return WriteInfoScalar(conn, SQL_CB_NULL, infoValuePtr, stringLengthPtr); + case SQL_NULL_COLLATION: + return WriteInfoScalar(conn, SQL_NC_HIGH, infoValuePtr, stringLengthPtr); + case SQL_MAX_CURSOR_NAME_LEN: + return WriteInfoScalar(conn, 128, infoValuePtr, stringLengthPtr); + + // DBMS Information + case SQL_DBMS_NAME: + return WriteInfoString(conn, "YDB", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_DBMS_VER: + return WriteInfoString(conn, conn->GetDbmsVersion().c_str(), infoValuePtr, bufferLength, stringLengthPtr); + + // Identifier Handling + case SQL_IDENTIFIER_QUOTE_CHAR: + return WriteInfoString(conn, "`", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_IDENTIFIER_CASE: + return WriteInfoScalar(conn, SQL_IC_SENSITIVE, infoValuePtr, stringLengthPtr); + + // Catalog Support + case SQL_CATALOG_NAME: + return WriteInfoString(conn, "Y", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_CATALOG_NAME_SEPARATOR: + return WriteInfoString(conn, "/", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_CATALOG_TERM: + return WriteInfoString(conn, "path", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_CATALOG_USAGE: + return WriteInfoScalar(conn, SQL_CU_DML_STATEMENTS, infoValuePtr, stringLengthPtr); + + // Schema Support (YDB doesn't use schemas) + case SQL_SCHEMA_USAGE: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + case SQL_SCHEMA_TERM: + return WriteInfoString(conn, "", infoValuePtr, bufferLength, stringLengthPtr); + + // Data Source Capabilities + case SQL_DATA_SOURCE_READ_ONLY: + return WriteInfoString( + conn, conn->IsDataSourceReadOnly() ? "Y" : "N", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_DATA_SOURCE_NAME: + return WriteInfoString(conn, conn->GetDataSourceName().c_str(), infoValuePtr, bufferLength, stringLengthPtr); + + // Result Set Capabilities + case SQL_MULT_RESULT_SETS: + return WriteInfoString(conn, "N", infoValuePtr, bufferLength, stringLengthPtr); + case SQL_DYNAMIC_CURSOR_ATTRIBUTES1: + case SQL_FORWARD_ONLY_CURSOR_ATTRIBUTES1: + case SQL_STATIC_CURSOR_ATTRIBUTES1: + return WriteInfoScalar(conn, SQL_CA1_NEXT, infoValuePtr, stringLengthPtr); + case SQL_CURSOR_COMMIT_BEHAVIOR: + case SQL_CURSOR_ROLLBACK_BEHAVIOR: + return WriteInfoScalar(conn, SQL_CB_CLOSE, infoValuePtr, stringLengthPtr); + + // Transaction Support + case SQL_TXN_CAPABLE: + return WriteInfoScalar(conn, SQL_TC_DML, infoValuePtr, stringLengthPtr); + case SQL_DEFAULT_TXN_ISOLATION: + return WriteInfoScalar(conn, SQL_TXN_SERIALIZABLE, infoValuePtr, stringLengthPtr); + case SQL_TXN_ISOLATION_OPTION: + return WriteInfoScalar( + conn, conn->GetSupportedTxnIsolationOptions(), infoValuePtr, stringLengthPtr); + + // Stored Procedures (not supported) + case SQL_PROCEDURES: + return WriteInfoString(conn, "N", infoValuePtr, bufferLength, stringLengthPtr); + + case SQL_OUTER_JOINS: + return WriteInfoString(conn, "Y", infoValuePtr, bufferLength, stringLengthPtr); + + // Positioned Operations (not supported) + case SQL_POSITIONED_STATEMENTS: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + + // Batch Operations (not supported) + case SQL_BATCH_SUPPORT: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + case SQL_BATCH_ROW_COUNT: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + case SQL_PARAM_ARRAY_ROW_COUNTS: + return WriteInfoScalar(conn, SQL_PARC_NO_BATCH, infoValuePtr, stringLengthPtr); + case SQL_PARAM_ARRAY_SELECTS: + return WriteInfoScalar(conn, SQL_PAS_NO_SELECT, infoValuePtr, stringLengthPtr); + + // Bookmarks (not supported) + case SQL_BOOKMARK_PERSISTENCE: + return WriteInfoScalar(conn, 0, infoValuePtr, stringLengthPtr); + + // Named Cursors (not supported) + case SQL_FILE_USAGE: + return WriteInfoScalar(conn, SQL_FILE_NOT_SUPPORTED, infoValuePtr, stringLengthPtr); + + // GetData Extensions + case SQL_GETDATA_EXTENSIONS: + return WriteInfoScalar(conn, SQL_GD_ANY_COLUMN | SQL_GD_ANY_ORDER, infoValuePtr, stringLengthPtr); + + // Async Execution (not supported) + case SQL_ASYNC_MODE: + return WriteInfoScalar(conn, SQL_AM_NONE, infoValuePtr, stringLengthPtr); + + case SQL_QUOTED_IDENTIFIER_CASE: + return WriteInfoScalar(conn, SQL_IC_SENSITIVE, infoValuePtr, stringLengthPtr); + + default: + return conn->AddError("HYC00", 0, "Optional feature not implemented"); + } +} + + +SQLRETURN NMetadata::GetFunctions(SQLUSMALLINT functionId, SQLUSMALLINT* supportedPtr) { + if (!supportedPtr) { + return SQL_ERROR; + } + + if (functionId == SQL_API_ALL_FUNCTIONS) { + std::memset(supportedPtr, 0, 100 * sizeof(SQLUSMALLINT)); + for (SQLUSMALLINT id = 0; id < 100; ++id) { + if (IsSupportedFunction(id)) { + supportedPtr[id] = SQL_TRUE; + } + } + return SQL_SUCCESS; + } + + if (functionId == SQL_API_ODBC3_ALL_FUNCTIONS) { + std::memset(supportedPtr, 0, SQL_API_ODBC3_ALL_FUNCTIONS_SIZE * sizeof(SQLUSMALLINT)); + for (SQLUSMALLINT id = 0; id < SQL_API_ODBC3_ALL_FUNCTIONS_SIZE * 16; ++id) { + if (IsSupportedFunction(id)) { + supportedPtr[id >> 4] |= (1 << (id & 0x000F)); + } + } + return SQL_SUCCESS; + } + + *supportedPtr = IsSupportedFunction(functionId) ? SQL_TRUE : SQL_FALSE; + return SQL_SUCCESS; +} + +SQLRETURN NMetadata::DescribeCol( + TStatement* stmt, + SQLUSMALLINT columnNumber, + SQLCHAR* columnName, + SQLSMALLINT bufferLength, + SQLSMALLINT* nameLengthPtr, + SQLSMALLINT* dataTypePtr, + SQLULEN* columnSizePtr, + SQLSMALLINT* decimalDigitsPtr, + SQLSMALLINT* nullablePtr) { + const auto& columns = stmt->GetColumnMeta(); + if (columnNumber < 1 || columnNumber > columns.size()) { + throw TOdbcException("07009", 0, "Invalid descriptor index"); + } + + const auto& column = columns[columnNumber - 1]; + const SQLRETURN nameRc = Diag::WriteOdbcString(*stmt, column.Name, columnName, bufferLength, nameLengthPtr); + if (nameRc != SQL_SUCCESS) { + return nameRc; + } + if (dataTypePtr) { + *dataTypePtr = column.SqlType; + } + if (columnSizePtr) { + *columnSizePtr = column.Size; + } + if (decimalDigitsPtr) { + *decimalDigitsPtr = column.DecimalDigits; + } + if (nullablePtr) { + *nullablePtr = column.Nullable; + } + return SQL_SUCCESS; +} + +SQLRETURN NMetadata::ColAttribute( + TStatement* stmt, + SQLUSMALLINT columnNumber, + SQLUSMALLINT fieldIdentifier, + SQLPOINTER characterAttributePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthAttributePtr, + SQLLEN* numericAttributePtr) { + SQLCHAR name[256] = {}; + SQLSMALLINT nameLength = 0; + SQLSMALLINT dataType = 0; + SQLULEN columnSize = 0; + SQLSMALLINT decimalDigits = 0; + SQLSMALLINT nullable = 0; + + const SQLRETURN describeRc = DescribeCol( + stmt, columnNumber, name, sizeof(name), &nameLength, &dataType, &columnSize, &decimalDigits, &nullable); + if (describeRc != SQL_SUCCESS) { + return describeRc; + } + + const auto setNumericAttr = [&](SQLLEN value) -> SQLRETURN { + if (!numericAttributePtr) { + return stmt->AddError("HY009", 0, "Invalid use of null pointer"); + } + *numericAttributePtr = value; + return SQL_SUCCESS; + }; + + switch (fieldIdentifier) { + case SQL_DESC_NAME: + case SQL_COLUMN_NAME: { + if (!characterAttributePtr && bufferLength != 0) { + return stmt->AddError("HY090", 0, "Invalid string or buffer length"); + } + const SQLSMALLINT fullLen = nameLength; + if (stringLengthAttributePtr) { + *stringLengthAttributePtr = fullLen; + } + if (bufferLength == 0) { + return fullLen == 0 ? SQL_SUCCESS + : stmt->AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + } + auto* out = reinterpret_cast(characterAttributePtr); + const SQLSMALLINT copyLen = static_cast(std::min(fullLen, bufferLength - 1)); + if (copyLen > 0) { + std::memcpy(out, name, static_cast(copyLen)); + } + if (out) { + out[copyLen] = '\0'; + } + if (copyLen < fullLen) { + return stmt->AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + } + return SQL_SUCCESS; + } + case SQL_DESC_TYPE: + case SQL_COLUMN_TYPE: + return setNumericAttr(dataType); + case SQL_DESC_LENGTH: + case SQL_COLUMN_LENGTH: + return setNumericAttr(static_cast(columnSize)); + case SQL_DESC_PRECISION: + case SQL_COLUMN_PRECISION: + return setNumericAttr(static_cast(columnSize)); + case SQL_DESC_SCALE: + case SQL_COLUMN_SCALE: + return setNumericAttr(decimalDigits); + case SQL_DESC_NULLABLE: + case SQL_COLUMN_NULLABLE: + return setNumericAttr(nullable); + default: + return stmt->AddError("HYC00", 0, "Optional feature not implemented"); + } +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/metadata.h b/odbc/src/metadata.h new file mode 100644 index 00000000000..59b83896c4b --- /dev/null +++ b/odbc/src/metadata.h @@ -0,0 +1,41 @@ +#pragma once + +#include "connection.h" +#include "statement.h" + +namespace NYdb::NOdbc { +namespace NMetadata { + +SQLRETURN GetInfo( + TConnection* conn, + SQLUSMALLINT infoType, + SQLPOINTER infoValuePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr); + +SQLRETURN GetFunctions( + SQLUSMALLINT functionId, + SQLUSMALLINT* supportedPtr); + +SQLRETURN DescribeCol( + TStatement* stmt, + SQLUSMALLINT columnNumber, + SQLCHAR* columnName, + SQLSMALLINT bufferLength, + SQLSMALLINT* nameLengthPtr, + SQLSMALLINT* dataTypePtr, + SQLULEN* columnSizePtr, + SQLSMALLINT* decimalDigitsPtr, + SQLSMALLINT* nullablePtr); + +SQLRETURN ColAttribute( + TStatement* stmt, + SQLUSMALLINT columnNumber, + SQLUSMALLINT fieldIdentifier, + SQLPOINTER characterAttributePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthAttributePtr, + SQLLEN* numericAttributePtr); + +} // namespace NMetadata +} // namespace NYdb::NOdbc diff --git a/odbc/src/odbc_driver.cpp b/odbc/src/odbc_driver.cpp new file mode 100644 index 00000000000..ec0564688dc --- /dev/null +++ b/odbc/src/odbc_driver.cpp @@ -0,0 +1,698 @@ +#include "environment.h" +#include "connection.h" +#include "statement.h" +#include "metadata.h" +#include "descriptor.h" + +#include "utils/util.h" +#include "utils/error_manager.h" + +#include +#include + +namespace { + template + Handle* GetHandle(SQLHANDLE handle) { + if (!handle) { + throw NYdb::NOdbc::TOdbcException("HY000", 0, "Invalid handle", SQL_INVALID_HANDLE); + } + return static_cast(handle); + } + +} + +extern "C" { + +SQLRETURN SQL_API SQLAllocHandle(SQLSMALLINT handleType, + SQLHANDLE inputHandle, + SQLHANDLE* outputHandle) { + if (!outputHandle) { + return SQL_INVALID_HANDLE; + } + + switch (handleType) { + case SQL_HANDLE_ENV: { + return NYdb::NOdbc::HandleOdbcExceptions( + inputHandle, + [&]() { + auto* const env = new NYdb::NOdbc::TEnvironment(); + *outputHandle = env; + env->SetLastReturnCode(SQL_SUCCESS); + return SQL_SUCCESS; + }, + NYdb::NOdbc::ENullInputHandlePolicy::Allow); + } + + case SQL_HANDLE_DBC: { + return NYdb::NOdbc::HandleOdbcExceptions(inputHandle, [&](auto* env) { + auto conn = std::make_unique(); + conn->SetEnvironment(env); + env->RegisterConnection(conn.get()); + auto* const raw = conn.release(); + *outputHandle = raw; + raw->SetLastReturnCode(SQL_SUCCESS); + return SQL_SUCCESS; + }); + } + case SQL_HANDLE_STMT: { + return NYdb::NOdbc::HandleOdbcExceptions(inputHandle, [&](auto* conn) { + auto stmt = conn->CreateStatement(); + auto* const raw = stmt.release(); + *outputHandle = raw; + raw->SetLastReturnCode(SQL_SUCCESS); + return SQL_SUCCESS; + }); + } + case SQL_HANDLE_DESC: { + return NYdb::NOdbc::HandleOdbcExceptions( + inputHandle, + [&](auto* conn) { + auto* const desc = new NYdb::NOdbc::TDescriptor( + NYdb::NOdbc::EDescType::Explicit, conn); + *outputHandle = desc; + desc->SetLastReturnCode(SQL_SUCCESS); + return SQL_SUCCESS; + }); + } + default: + return SQL_ERROR; + } +} + +SQLRETURN SQL_API SQLFreeHandle(SQLSMALLINT handleType, SQLHANDLE handle) { + switch (handleType) { + case SQL_HANDLE_ENV: { + return NYdb::NOdbc::HandleOdbcExceptionsConsuming(handle, [](auto* env) { + if (!env->GetConnectionsSnapshot().empty()) { + return env->AddError("HY010", 0, "Connection handles are still allocated"); + } + delete env; + return static_cast(SQL_SUCCESS); + }); + } + case SQL_HANDLE_DBC: { + return NYdb::NOdbc::HandleOdbcExceptionsConsuming(handle, [](auto* conn) { + if (conn->HasChildren()) { + return conn->AddError("HY010", 0, "Statement or descriptor handles are still allocated"); + } + auto* env = conn->GetEnvironment(); + if (env != nullptr){ + env->UnregisterConnection(conn); + } + delete conn; + return static_cast(SQL_SUCCESS); + }); + } + case SQL_HANDLE_STMT: { + return NYdb::NOdbc::HandleOdbcExceptionsConsuming(handle, [](auto* stmt) { + delete stmt; + return SQL_SUCCESS; + }); + } + case SQL_HANDLE_DESC: { + return NYdb::NOdbc::HandleOdbcExceptionsConsuming(handle, [](auto* desc) { + if (desc->GetDescType() != NYdb::NOdbc::EDescType::Explicit) { + return desc->AddError( + "HY017", 0, "Invalid use of an automatically allocated descriptor handle"); + } + delete desc; + return static_cast(SQL_SUCCESS); + }); + } + default: + return SQL_ERROR; + } +} + +SQLRETURN SQL_API SQLSetEnvAttr(SQLHENV environmentHandle, + SQLINTEGER attribute, + SQLPOINTER value, + SQLINTEGER stringLength) { + auto env = static_cast(environmentHandle); + if (!env) { + return SQL_INVALID_HANDLE; + } + + return NYdb::NOdbc::HandleOdbcExceptions(env, [&]() { + return env->SetAttribute(attribute, value, stringLength); + }); +} + +SQLRETURN SQL_API SQLGetEnvAttr(SQLHENV environmentHandle, + SQLINTEGER attribute, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr) { + auto env = static_cast(environmentHandle); + if (!env) { + return SQL_INVALID_HANDLE; + } + return NYdb::NOdbc::HandleOdbcExceptions(env, [&]() { + return env->GetAttribute(attribute, value, bufferLength, stringLengthPtr); + }); +} + +SQLRETURN SQL_API SQLDriverConnect(SQLHDBC connectionHandle, + SQLHWND /*WindowHandle*/, + SQLCHAR* inConnectionString, + SQLSMALLINT stringLength1, + SQLCHAR* /*outConnectionString*/, + SQLSMALLINT /*bufferLength*/, + SQLSMALLINT* /*stringLength2Ptr*/, + SQLUSMALLINT /*driverCompletion*/) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return conn->DriverConnect(NYdb::NOdbc::GetString(inConnectionString, stringLength1)); + }); +} + +SQLRETURN SQL_API SQLConnect(SQLHDBC connectionHandle, + SQLCHAR* serverName, SQLSMALLINT nameLength1, + SQLCHAR* userName, SQLSMALLINT nameLength2, + SQLCHAR* authentication, SQLSMALLINT nameLength3) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return conn->Connect(NYdb::NOdbc::GetString(serverName, nameLength1), + NYdb::NOdbc::GetString(userName, nameLength2), + NYdb::NOdbc::GetString(authentication, nameLength3)); + }); +} + +SQLRETURN SQL_API SQLDisconnect(SQLHDBC connectionHandle) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return conn->Disconnect(); + }); +} + +SQLRETURN SQL_API SQLExecDirect(SQLHSTMT statementHandle, + SQLCHAR* statementText, + SQLINTEGER textLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + auto ret = stmt->Prepare(NYdb::NOdbc::GetString(statementText, textLength)); + if (ret != SQL_SUCCESS) { + return ret; + } + return stmt->Execute(); + }); +} + +SQLRETURN SQL_API SQLExecDirectW(SQLHSTMT statementHandle, + SQLWCHAR* statementText, + SQLINTEGER textLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + auto ret = stmt->Prepare(NYdb::NOdbc::GetString(statementText, textLength)); + if (ret != SQL_SUCCESS) { + return ret; + } + return stmt->Execute(); + }); +} + +SQLRETURN SQL_API SQLPrepare(SQLHSTMT statementHandle, + SQLCHAR* statementText, + SQLINTEGER textLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Prepare(NYdb::NOdbc::GetString(statementText, textLength)); + }); +} + +SQLRETURN SQL_API SQLPrepareW(SQLHSTMT statementHandle, + SQLWCHAR* statementText, + SQLINTEGER textLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Prepare(NYdb::NOdbc::GetString(statementText, textLength)); + }); +} + +SQLRETURN SQL_API SQLExecute(SQLHSTMT statementHandle) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Execute(); + }); +} + +SQLRETURN SQL_API SQLFetch(SQLHSTMT statementHandle) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Fetch(); + }); +} + +SQLRETURN SQL_API SQLGetData(SQLHSTMT statementHandle, + SQLUSMALLINT columnNumber, + SQLSMALLINT targetType, + SQLPOINTER targetValue, + SQLLEN bufferLength, + SQLLEN* strLenOrInd) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->GetData(columnNumber, targetType, targetValue, bufferLength, strLenOrInd); + }); +} + +SQLRETURN SQL_API SQLBindCol(SQLHSTMT statementHandle, + SQLUSMALLINT columnNumber, + SQLSMALLINT targetType, + SQLPOINTER targetValue, + SQLLEN bufferLength, + SQLLEN* strLenOrInd) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->BindCol(columnNumber, targetType, targetValue, bufferLength, strLenOrInd); + }); +} + +SQLRETURN SQL_API SQLGetDiagRec(SQLSMALLINT handleType, + SQLHANDLE handle, + SQLSMALLINT recNumber, + SQLCHAR* sqlState, + SQLINTEGER* nativeError, + SQLCHAR* messageText, + SQLSMALLINT bufferLength, + SQLSMALLINT* textLength) { + switch (handleType) { + case SQL_HANDLE_ENV: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* env) { + return env->GetDiagRec(recNumber, sqlState, nativeError, messageText, bufferLength, textLength); + }); + } + case SQL_HANDLE_DBC: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* conn) { + return conn->GetDiagRec(recNumber, sqlState, nativeError, messageText, bufferLength, textLength); + }); + } + case SQL_HANDLE_STMT: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* stmt) { + return stmt->GetDiagRec(recNumber, sqlState, nativeError, messageText, bufferLength, textLength); + }); + } + case SQL_HANDLE_DESC: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* desc) { + return desc->GetDiagRec(recNumber, sqlState, nativeError, messageText, bufferLength, textLength); + }); + } + default: + return SQL_ERROR; + } +} + +SQLRETURN SQL_API SQLGetDiagField(SQLSMALLINT handleType, + SQLHANDLE handle, + SQLSMALLINT recNumber, + SQLSMALLINT diagIdentifier, + SQLPOINTER diagInfoPtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) { + switch (handleType) { + case SQL_HANDLE_ENV: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* env) { + return env->GetDiagField(recNumber, diagIdentifier, diagInfoPtr, bufferLength, stringLengthPtr); + }); + } + case SQL_HANDLE_DBC: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* conn) { + return conn->GetDiagField(recNumber, diagIdentifier, diagInfoPtr, bufferLength, stringLengthPtr); + }); + } + case SQL_HANDLE_STMT: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* stmt) { + return stmt->GetDiagField(recNumber, diagIdentifier, diagInfoPtr, bufferLength, stringLengthPtr); + }); + } + case SQL_HANDLE_DESC: { + return NYdb::NOdbc::HandleOdbcDiagnostics(handle, [&](auto* desc) { + return desc->GetDiagField(recNumber, diagIdentifier, diagInfoPtr, bufferLength, stringLengthPtr); + }); + } + default: + return SQL_ERROR; + } +} + +SQLRETURN SQL_API SQLBindParameter(SQLHSTMT statementHandle, + SQLUSMALLINT paramNumber, + SQLSMALLINT inputOutputType, + SQLSMALLINT valueType, + SQLSMALLINT parameterType, + SQLULEN columnSize, + SQLSMALLINT decimalDigits, + SQLPOINTER parameterValuePtr, + SQLLEN bufferLength, + SQLLEN* strLenOrIndPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->BindParameter(paramNumber, inputOutputType, valueType, parameterType, columnSize, decimalDigits, parameterValuePtr, bufferLength, strLenOrIndPtr); + }); +} + +SQLRETURN SQL_API SQLEndTran(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT completionType) { + switch (handleType) { + case SQL_HANDLE_DBC: { + return NYdb::NOdbc::HandleOdbcExceptions(handle, [&](auto* conn) { + if (completionType == SQL_COMMIT) { + return conn->CommitTx(); + } else if (completionType == SQL_ROLLBACK) { + return conn->RollbackTx(); + } else { + throw NYdb::NOdbc::TOdbcException("HY012", 0, "Invalid transaction operation code"); + } + }); + } + case SQL_HANDLE_ENV: { + return NYdb::NOdbc::HandleOdbcExceptions(handle, [&](auto* env) -> SQLRETURN { + return env->EndTran(completionType); + }); + } + default: + return SQL_INVALID_HANDLE; + } +} + +SQLRETURN SQL_API SQLSetConnectAttr(SQLHDBC connectionHandle, SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER stringLength) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return conn->SetConnectAttr(attribute, value, stringLength); + }); +} + +SQLRETURN SQL_API SQLGetConnectAttr(SQLHDBC connectionHandle, SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return conn->GetConnectAttr(attribute, value, bufferLength, stringLengthPtr); + }); +} + +SQLRETURN SQL_API SQLColumns(SQLHSTMT statementHandle, + SQLCHAR* catalogName, SQLSMALLINT nameLength1, + SQLCHAR* schemaName, SQLSMALLINT nameLength2, + SQLCHAR* tableName, SQLSMALLINT nameLength3, + SQLCHAR* columnName, SQLSMALLINT nameLength4) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Columns( + NYdb::NOdbc::GetString(catalogName, nameLength1), + NYdb::NOdbc::GetString(schemaName, nameLength2), + NYdb::NOdbc::GetString(tableName, nameLength3), + NYdb::NOdbc::GetString(columnName, nameLength4)); + }); +} + +SQLRETURN SQL_API SQLTables(SQLHSTMT statementHandle, + SQLCHAR* catalogName, SQLSMALLINT nameLength1, + SQLCHAR* schemaName, SQLSMALLINT nameLength2, + SQLCHAR* tableName, SQLSMALLINT nameLength3, + SQLCHAR* tableType, SQLSMALLINT nameLength4) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Tables( + NYdb::NOdbc::GetString(catalogName, nameLength1), + NYdb::NOdbc::GetString(schemaName, nameLength2), + NYdb::NOdbc::GetString(tableName, nameLength3), + NYdb::NOdbc::GetString(tableType, nameLength4)); + }); +} + +SQLRETURN SQL_API SQLCloseCursor(SQLHSTMT statementHandle) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Close(false); + }); +} + +SQLRETURN SQL_API SQLFreeStmt(SQLHSTMT statementHandle, SQLUSMALLINT option) { + if (option == SQL_DROP) { + return SQLFreeHandle(SQL_HANDLE_STMT, statementHandle); + } + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) -> SQLRETURN { + switch (option) { + case SQL_CLOSE: + return stmt->Close(true); + case SQL_UNBIND: + stmt->UnbindColumns(); + return SQL_SUCCESS; + case SQL_RESET_PARAMS: + stmt->ResetParams(); + return SQL_SUCCESS; + default: + throw NYdb::NOdbc::TOdbcException("HY000", 0, "Invalid option"); + } + }); +} + +SQLRETURN SQL_API SQLFetchScroll(SQLHSTMT statementHandle, SQLSMALLINT fetchOrientation, SQLLEN fetchOffset) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + if (fetchOrientation == SQL_FETCH_NEXT) { + return stmt->Fetch(); + } else { + throw NYdb::NOdbc::TOdbcException("HYC00", 0, "Only SQL_FETCH_NEXT is supported"); + } + //TODO other fetch-orientation + }); +} + +SQLRETURN SQL_API SQLRowCount(SQLHSTMT statementHandle, SQLLEN* rowCount) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->RowCount(rowCount); + }); +} + +SQLRETURN SQL_API SQLNumResultCols(SQLHSTMT statementHandle, SQLSMALLINT* colCount) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->NumResultCols(colCount); + }); +} + +SQLRETURN SQL_API SQLDescribeCol( + SQLHSTMT statementHandle, + SQLUSMALLINT columnNumber, + SQLCHAR* columnName, + SQLSMALLINT bufferLength, + SQLSMALLINT* nameLengthPtr, + SQLSMALLINT* dataTypePtr, + SQLULEN* columnSizePtr, + SQLSMALLINT* decimalDigitsPtr, + SQLSMALLINT* nullablePtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return NYdb::NOdbc::NMetadata::DescribeCol( + stmt, + columnNumber, + columnName, + bufferLength, + nameLengthPtr, + dataTypePtr, + columnSizePtr, + decimalDigitsPtr, + nullablePtr); + }); +} + +SQLRETURN SQL_API SQLMoreResults(SQLHSTMT) { + // YDB ODBC currently exposes only one result set per statement. + return SQL_NO_DATA; +} + +SQLRETURN SQL_API SQLGetFunctions(SQLHDBC connectionHandle, SQLUSMALLINT functionId, SQLUSMALLINT* supportedPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto*) { + return NYdb::NOdbc::NMetadata::GetFunctions(functionId, supportedPtr); + }); +} + +SQLRETURN SQL_API SQLSetStmtAttr(SQLHSTMT statementHandle, SQLINTEGER attribute, SQLPOINTER value, SQLINTEGER stringLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->SetStmtAttr(attribute, value, stringLength); + }); +} + +SQLRETURN SQL_API SQLGetStmtAttr( + SQLHSTMT statementHandle, + SQLINTEGER attribute, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->GetStmtAttr(attribute, value, bufferLength, stringLengthPtr); + }); +} + +SQLRETURN SQL_API SQLGetInfo(SQLHDBC connectionHandle, + SQLUSMALLINT infoType, + SQLPOINTER infoValuePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + return NYdb::NOdbc::NMetadata::GetInfo(conn, infoType, infoValuePtr, bufferLength, stringLengthPtr); + }); +} + +SQLRETURN SQL_API SQLGetTypeInfo(SQLHSTMT statementHandle, SQLSMALLINT dataType) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->GetTypeInfo(dataType); + }); +} + +SQLRETURN SQL_API SQLStatistics(SQLHSTMT statementHandle, + SQLCHAR* catalogName, SQLSMALLINT nameLength1, + SQLCHAR* schemaName, SQLSMALLINT nameLength2, + SQLCHAR* tableName, SQLSMALLINT nameLength3, + SQLUSMALLINT unique, SQLUSMALLINT reserved) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Statistics( + NYdb::NOdbc::GetString(catalogName, nameLength1), + NYdb::NOdbc::GetString(schemaName, nameLength2), + NYdb::NOdbc::GetString(tableName, nameLength3), + unique, + reserved); + }); +} + +SQLRETURN SQL_API SQLSpecialColumns(SQLHSTMT statementHandle, + SQLUSMALLINT identifierType, + SQLCHAR* catalogName, SQLSMALLINT nameLength1, + SQLCHAR* schemaName, SQLSMALLINT nameLength2, + SQLCHAR* tableName, SQLSMALLINT nameLength3, + SQLUSMALLINT scope, + SQLUSMALLINT nullable) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + (void)nullable; + return stmt->SpecialColumns( + NYdb::NOdbc::GetString(catalogName, nameLength1), + NYdb::NOdbc::GetString(schemaName, nameLength2), + NYdb::NOdbc::GetString(tableName, nameLength3), + identifierType, + scope); + }); +} + +SQLRETURN SQL_API SQLColAttribute(SQLHSTMT statementHandle, + SQLUSMALLINT columnNumber, + SQLUSMALLINT fieldIdentifier, + SQLPOINTER characterAttributePtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthAttributePtr, + SQLLEN* numericAttributePtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return NYdb::NOdbc::NMetadata::ColAttribute( + stmt, columnNumber, fieldIdentifier, characterAttributePtr, bufferLength, + stringLengthAttributePtr, numericAttributePtr); + }); +} + +SQLRETURN SQL_API SQLNumParams(SQLHSTMT statementHandle, SQLSMALLINT* paramCountPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->NumParams(paramCountPtr); + }); +} + +SQLRETURN SQL_API SQLDescribeParam(SQLHSTMT statementHandle, SQLUSMALLINT paramNumber, SQLSMALLINT* dataTypePtr, + SQLULEN* paramSizePtr, SQLSMALLINT* decimalDigitsPtr, SQLSMALLINT* nullablePtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->DescribeParam(paramNumber, dataTypePtr, paramSizePtr, decimalDigitsPtr, nullablePtr); + }); +} + +SQLRETURN SQL_API SQLParamData(SQLHSTMT statementHandle, SQLPOINTER* valuePtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->ParamData(valuePtr); + }); +} + +SQLRETURN SQL_API SQLPutData(SQLHSTMT statementHandle, SQLPOINTER data, SQLLEN strLenOrInd) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->PutData(data, strLenOrInd); + }); +} + +SQLRETURN SQL_API SQLCancel(SQLHSTMT statementHandle) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->Cancel(); + }); +} + +SQLRETURN SQL_API SQLNativeSql(SQLHDBC connectionHandle, + SQLCHAR* inNativeSql, + SQLINTEGER textLength1, + SQLCHAR* outNativeSql, + SQLINTEGER bufferLength, + SQLINTEGER* outLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(connectionHandle, [&](auto* conn) { + const std::string inSql = textLength1 == SQL_NTS + ? reinterpret_cast(inNativeSql) + : NYdb::NOdbc::GetString(inNativeSql, static_cast(textLength1)); + return conn->NativeSql(inSql, outNativeSql, bufferLength, outLengthPtr); + }); +} + +SQLRETURN SQL_API SQLSetCursorName(SQLHSTMT statementHandle, SQLCHAR* cursorName, SQLSMALLINT nameLength) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->SetCursorName(NYdb::NOdbc::GetString(cursorName, nameLength)); + }); +} + +SQLRETURN SQL_API SQLGetCursorName(SQLHSTMT statementHandle, + SQLCHAR* cursorName, + SQLSMALLINT bufferLength, + SQLSMALLINT* nameLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->GetCursorName(cursorName, bufferLength, nameLengthPtr); + }); +} + +SQLRETURN SQL_API SQLPrimaryKeys(SQLHSTMT statementHandle, + SQLCHAR* catalogName, SQLSMALLINT nameLength1, + SQLCHAR* schemaName, SQLSMALLINT nameLength2, + SQLCHAR* tableName, SQLSMALLINT nameLength3) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->PrimaryKeys( + NYdb::NOdbc::GetString(catalogName, nameLength1), + NYdb::NOdbc::GetString(schemaName, nameLength2), + NYdb::NOdbc::GetString(tableName, nameLength3)); + }); +} + +SQLRETURN SQL_API SQLForeignKeys(SQLHSTMT statementHandle, + SQLCHAR* pkCatalogName, SQLSMALLINT nameLength1, + SQLCHAR* pkSchemaName, SQLSMALLINT nameLength2, + SQLCHAR* pkTableName, SQLSMALLINT nameLength3, + SQLCHAR* fkCatalogName, SQLSMALLINT nameLength4, + SQLCHAR* fkSchemaName, SQLSMALLINT nameLength5, + SQLCHAR* fkTableName, SQLSMALLINT nameLength6) { + return NYdb::NOdbc::HandleOdbcExceptions(statementHandle, [&](auto* stmt) { + return stmt->ForeignKeys( + NYdb::NOdbc::GetString(pkCatalogName, nameLength1), + NYdb::NOdbc::GetString(pkSchemaName, nameLength2), + NYdb::NOdbc::GetString(pkTableName, nameLength3), + NYdb::NOdbc::GetString(fkCatalogName, nameLength4), + NYdb::NOdbc::GetString(fkSchemaName, nameLength5), + NYdb::NOdbc::GetString(fkTableName, nameLength6)); + }); +} + +SQLRETURN SQL_API SQLGetDescField(SQLHDESC descriptorHandle, SQLSMALLINT recNumber, SQLSMALLINT fieldIdentifier, + SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(descriptorHandle, [&](auto* desc) { + return desc->GetDescField(recNumber, fieldIdentifier, value, bufferLength, stringLengthPtr); + }); +} + +SQLRETURN SQL_API SQLGetDescRec(SQLHDESC descriptorHandle, SQLSMALLINT recNumber, SQLCHAR* name, + SQLSMALLINT bufferLength, SQLSMALLINT* stringLengthPtr, SQLSMALLINT* typePtr, + SQLSMALLINT* subTypePtr, SQLLEN* lengthPtr, SQLSMALLINT* precisionPtr, + SQLSMALLINT* scalePtr, SQLSMALLINT* nullablePtr) { + return NYdb::NOdbc::HandleOdbcExceptions(descriptorHandle, [&](auto* desc) { + return desc->GetDescRec(recNumber, name, bufferLength, stringLengthPtr, typePtr, subTypePtr, + lengthPtr, precisionPtr, scalePtr, nullablePtr); + }); +} + +SQLRETURN SQL_API SQLSetDescField(SQLHDESC descriptorHandle, SQLSMALLINT recNumber, SQLSMALLINT fieldIdentifier, + SQLPOINTER value, SQLINTEGER bufferLength) { + return NYdb::NOdbc::HandleOdbcExceptions(descriptorHandle, [&](auto* desc) { + return desc->SetDescField(recNumber, fieldIdentifier, value, bufferLength); + }); +} + +SQLRETURN SQL_API SQLSetDescRec(SQLHDESC descriptorHandle, SQLSMALLINT recNumber, SQLSMALLINT type, + SQLSMALLINT subType, SQLLEN length, SQLSMALLINT precision, SQLSMALLINT scale, + SQLPOINTER dataPtr, SQLLEN* stringLengthPtr, SQLLEN* indicatorPtr) { + return NYdb::NOdbc::HandleOdbcExceptions(descriptorHandle, [&](auto* desc) { + return desc->SetDescRec(recNumber, type, subType, length, precision, scale, dataPtr, + stringLengthPtr, indicatorPtr); + }); +} + +SQLRETURN SQL_API SQLCopyDesc(SQLHDESC sourceDesc, SQLHDESC targetDesc) { + return NYdb::NOdbc::HandleOdbcExceptions(sourceDesc, [&](auto* src) { + return src->CopyDesc(NYdb::NOdbc::TDescriptor::FromHandle(targetDesc)); + }); +} + +} diff --git a/odbc/src/statement.cpp b/odbc/src/statement.cpp new file mode 100644 index 00000000000..419d8f53783 --- /dev/null +++ b/odbc/src/statement.cpp @@ -0,0 +1,991 @@ +#include "statement.h" + +#include "utils/convert.h" +#include "utils/attr.h" +#include "utils/types.h" +#include "utils/diag.h" +#include "utils/escape.h" +#include "utils/param_rewrite.h" +#include "utils/sql_like.h" +#include "utils/type_info_rows.h" +#include "utils/util.h" +#include "utils/status_util.h" + +#include +#include +#include +#include + +#include + +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +namespace { + + size_t CTypeSize(SQLSMALLINT type, SQLLEN bufferLength) { + switch (type) { + case SQL_C_CHAR: case SQL_C_BINARY: return std::max(bufferLength, 0); + case SQL_C_BIT: case SQL_C_TINYINT: case SQL_C_UTINYINT: return sizeof(SQLCHAR); + case SQL_C_SHORT: case SQL_C_USHORT: return sizeof(SQLSMALLINT); + case SQL_C_LONG: case SQL_C_ULONG: return sizeof(SQLINTEGER); + case SQL_C_SBIGINT: case SQL_C_UBIGINT: return sizeof(SQLBIGINT); + case SQL_C_FLOAT: return sizeof(SQLREAL); + case SQL_C_DOUBLE: return sizeof(SQLDOUBLE); + case SQL_C_TYPE_DATE: return sizeof(SQL_DATE_STRUCT); + case SQL_C_TYPE_TIME: return sizeof(SQL_TIME_STRUCT); + case SQL_C_TYPE_TIMESTAMP: return sizeof(SQL_TIMESTAMP_STRUCT); + case SQL_C_GUID: return sizeof(SQLGUID); + default: + return static_cast(std::max(bufferLength, 0)); + } + } + + template + T* OffsetPointer(T* pointer, SQLULEN offset, SQLULEN row, SQLULEN stride) { + if (!pointer) { + return nullptr; + } + auto* bytes = reinterpret_cast(pointer); + return reinterpret_cast(bytes + offset + row * stride); + } + + TBoundParam ParamAt(const TBoundParam& param, SQLULEN row, SQLULEN bindType, SQLULEN offset) { + TBoundParam adjusted = param; + const SQLULEN dataStride = bindType == SQL_PARAM_BIND_BY_COLUMN + ? CTypeSize(param.ValueType, param.BufferLength) + : bindType; + const SQLULEN indicatorStride = bindType == SQL_PARAM_BIND_BY_COLUMN + ? sizeof(SQLLEN) + : bindType; + adjusted.ParameterValuePtr = OffsetPointer( + static_cast(param.ParameterValuePtr), offset, row, dataStride); + adjusted.StrLenOrIndPtr = OffsetPointer( + param.StrLenOrIndPtr, offset, row, indicatorStride); + return adjusted; + } + + bool StartsWithStatement( + std::string_view queryText, + std::initializer_list keywords) { + size_t i = 0; + while (i < queryText.size()) { + if (std::isspace(static_cast(queryText[i]))) { + ++i; + } else if (queryText[i] == '-' && i + 1 < queryText.size() && queryText[i + 1] == '-') { + while (i < queryText.size() && queryText[i] != '\n') { + ++i; + } + } else if (queryText[i] == '/' && i + 1 < queryText.size() && queryText[i + 1] == '*') { + i += 2; + while (i + 1 < queryText.size() && !(queryText[i] == '*' && queryText[i + 1] == '/')) { + ++i; + } + if (i + 1 < queryText.size()) { + i += 2; + } else { + i = queryText.size(); + } + } else { + break; + } + } + const size_t remaining = queryText.size() - i; + for (const std::string_view keyword : keywords) { + if (StartsWithPrefix( + queryText.data() + i, remaining, keyword.data(), keyword.size())) { + return true; + } + } + return false; + } + + std::optional ExtractAffectedRows(const NQuery::TExecuteQueryResult& result) { + const auto& stats = result.GetStats(); + if (!stats) { + return std::nullopt; + } + + const uint64_t maxSqlLen = static_cast(std::numeric_limits::max()); + uint64_t affectedRows = 0; + bool hasTableAccess = false; + for (const auto& phase : stats->GetQueryPhases()) { + for (const auto& table : phase.GetTableAccess()) { + hasTableAccess = true; + const uint64_t updates = table.GetUpdates().GetRows(); + const uint64_t deletes = table.GetDeletes().GetRows(); + if (updates > maxSqlLen - affectedRows) { + return std::nullopt; + } + affectedRows += updates; + if (deletes > maxSqlLen - affectedRows) { + return std::nullopt; + } + affectedRows += deletes; + } + } + if (affectedRows == 0 && (!hasTableAccess || !result.GetResultSets().empty())) { + return std::nullopt; + } + return static_cast(affectedRows); + } + +} + +TStatement::TStatement(TConnection* conn) + : Conn_(conn) + , AppRowDesc_(EDescType::AppRow, conn) + , AppParamDesc_(EDescType::AppParam, conn) + , ImpRowDesc_(EDescType::ImpRow, conn) + , ImpParamDesc_(EDescType::ImpParam, conn) + , CurrentAppRowDesc_(&AppRowDesc_) + , CurrentAppParamDesc_(&AppParamDesc_) { + Conn_->RegisterStatement(this); +} + +TStatement::~TStatement() { + CurrentAppRowDesc_->Detach(this); + CurrentAppParamDesc_->Detach(this); + Conn_->UnregisterStatement(this); +} + +void TStatement::DetachDescriptor(TDescriptor* desc) { + if (CurrentAppRowDesc_ == desc) { + CurrentAppRowDesc_ = &AppRowDesc_; + } + if (CurrentAppParamDesc_ == desc) { + CurrentAppParamDesc_ = &AppParamDesc_; + } + desc->Detach(this); +} + +SQLRETURN TStatement::Prepare(const std::string& statementText) { + RowsFetched_ = 0; + RowCount_ = -1; + SetCursor(nullptr); + PreparedQuery_ = statementText; + IsPrepared_ = true; + ParamCount_ = CountOdbcParams(PreparedQuery_); + while (ImpParamDesc_.GetRecordCount() > ParamCount_) { + ImpParamDesc_.RemoveRecord(ImpParamDesc_.GetRecordCount()); + } + for (SQLSMALLINT i = 1; i <= ParamCount_; ++i) { + TDescRecord& record = ImpParamDesc_.Record(i); + record.Nullable = SQL_NULLABLE_UNKNOWN; + } + return SQL_SUCCESS; +} + +SQLRETURN TStatement::Execute() { + if (!IsPrepared_ || PreparedQuery_.empty()) { + throw TOdbcException("HY007", 0, "No prepared statement"); + } + if (ParamCount_ > 0 && CurrentAppParamDesc_->GetArraySize() > 1 + && !StartsWithStatement(PreparedQuery_, {"INSERT", "UPDATE", "DELETE", "UPSERT", "REPLACE"})) { + return AddError("HYC00", 0, "Parameter arrays are supported only for data-modification statements"); + } + const SQLUSMALLINT next = FindNextNeedDataParam(); + if (next != 0) { + if (CurrentAppParamDesc_->GetArraySize() > 1) { + return AddError("HYC00", 0, "Data-at-execution parameter arrays are not supported"); + } + NeedDataParam_ = next; + InAtExec_ = true; + NeedDataTokenDelivered_ = false; + return SQL_NEED_DATA; + } + InAtExec_ = false; + NeedDataParam_ = 0; + return ExecuteInternal(); +} + +SQLRETURN TStatement::ExecuteInternal() { + RowCount_ = 0; + bool hasSuccessfulParamSet = false; + bool rowCountUsable = true; + const SQLULEN paramsetSize = ParamCount_ > 0 ? CurrentAppParamDesc_->GetArraySize() : 1; + SQLUSMALLINT* const operations = CurrentAppParamDesc_->GetArrayStatusPtr(); + SQLUSMALLINT* const statuses = ImpParamDesc_.GetArrayStatusPtr(); + SQLULEN* const processed = ImpParamDesc_.GetRowsProcessedPtr(); + if (processed) { + *processed = 0; + } + if (statuses) { + std::fill_n(statuses, paramsetSize, SQL_PARAM_UNUSED); + } + + SQLRETURN result = SQL_SUCCESS; + for (SQLULEN paramSet = 0; paramSet < paramsetSize; ++paramSet) { + if (operations && operations[paramSet] == SQL_PARAM_IGNORE) { + if (processed) { + *processed = paramSet + 1; + } + continue; + } + if (operations && operations[paramSet] != SQL_PARAM_PROCEED) { + if (statuses) { + statuses[paramSet] = SQL_PARAM_ERROR; + } + if (processed) { + *processed = paramSet + 1; + } + return AddError("HY024", 0, "Invalid parameter operation value"); + } + std::optional affectedRows; + SQLRETURN rc; + try { + rc = ExecuteParamSet(paramSet, affectedRows); + } catch (...) { + if (!hasSuccessfulParamSet) { + RowCount_ = -1; + } + throw; + } + if (statuses) { + statuses[paramSet] = rc == SQL_SUCCESS_WITH_INFO + ? SQL_PARAM_SUCCESS_WITH_INFO + : rc == SQL_SUCCESS ? SQL_PARAM_SUCCESS : SQL_PARAM_ERROR; + } + if (processed) { + *processed = paramSet + 1; + } + if (rc == SQL_ERROR) { + if (!hasSuccessfulParamSet) { + RowCount_ = -1; + } + return SQL_ERROR; + } + hasSuccessfulParamSet = true; + if (rowCountUsable) { + if (!affectedRows || *affectedRows > std::numeric_limits::max() - RowCount_) { + RowCount_ = -1; + rowCountUsable = false; + } else { + RowCount_ += *affectedRows; + } + } + if (rc == SQL_SUCCESS_WITH_INFO) { + result = SQL_SUCCESS_WITH_INFO; + } + } + return result; +} + +SQLRETURN TStatement::ExecuteParamSet( + SQLULEN paramSet, + std::optional& affectedRows) +{ + RowsFetched_ = 0; + SetCursor(nullptr); + auto client = Conn_->GetClient(); + if (!client) { + throw TOdbcException("HY000", 0, "No client connection"); + } + NYdb::TParams params = NYdb::TParamsBuilder().Build(); + const SQLRETURN buildRc = BuildParams(params, paramSet); + if (buildRc != SQL_SUCCESS) { + return buildRc; + } + + if (Conn_->GetAutocommit()) { + Conn_->ResetTx(); + Conn_->ResetQuerySession(); + const NYdb::NRetry::TRetryOperationSettings retrySettings = MakeAutocommitRetrySettings(); + + const NYdb::TStatus execStatus = client->RetryQuerySync( + [this, ¶ms, &affectedRows](NQuery::TSession session) -> NYdb::TStatus { + NQuery::TExecuteQueryResult result = ExecuteQuery(session, params); + if (!result.IsSuccess()) { + return StatusFrom(result); + } + affectedRows = ExtractAffectedRows(result); + SetCursor(CreateExecCursor(result)); + return NYdb::TStatus(EStatus::SUCCESS, NYdb::NIssue::TIssues()); + }, + retrySettings); + + NStatusHelpers::ThrowOnError(execStatus); + } else { + NQuery::TSession& session = Conn_->GetOrCreateQuerySession(); + NQuery::TExecuteQueryResult result = ExecuteQuery(session, params); + NStatusHelpers::ThrowOnError(result); + affectedRows = ExtractAffectedRows(result); + SetCursor(CreateExecCursor(result)); + } + InAtExec_ = false; + NeedDataParam_ = 0; + NeedDataTokenDelivered_ = false; + for (SQLSMALLINT i = 1; i <= CurrentAppParamDesc_->GetRecordCount(); ++i) { + if (TDescRecord* param = CurrentAppParamDesc_->FindRecord(i); param && param->AtExec) { + param->AtExecComplete = false; + param->AtExecChunk.clear(); + } + } + return SQL_SUCCESS; +} + +SQLUSMALLINT TStatement::FindNextNeedDataParam() const { + for (SQLSMALLINT i = 1; i <= CurrentAppParamDesc_->GetRecordCount(); ++i) { + const TDescRecord* record = CurrentAppParamDesc_->FindRecord(i); + if (record && record->AtExec && !record->AtExecComplete) { + return static_cast(i); + } + } + return 0; +} + +NYdb::NRetry::TRetryOperationSettings TStatement::MakeAutocommitRetrySettings() { + NYdb::NRetry::TRetryOperationSettings settings; + settings.Idempotent(false); + SQLUINTEGER queryTimeoutSec = Attributes_.GetQueryTimeoutSec(); + if (queryTimeoutSec > 0) { + const TDuration deadline = TDuration::Seconds(queryTimeoutSec); + settings.MaxTimeout(deadline).GetSessionClientTimeout(deadline); + } + return settings; +} + +NQuery::TExecuteQueryResult TStatement::ExecuteQuery( + NQuery::TSession& session, + const NYdb::TParams& params) +{ + const std::string sqlAfterEscapes = Attributes_.GetNoScanMode() == SQL_NOSCAN_ON + ? PreparedQuery_ + : RewriteOdbcEscapes(PreparedQuery_); + const std::vector activeParams = GetBoundParams(0); + const TParamRewriteResult rewritten = RewriteOdbcQuestionMarks(sqlAfterEscapes, activeParams); + if (!rewritten.Success) { + throw TOdbcException(rewritten.SqlState, 0, rewritten.Message); + } + const bool isDdl = StartsWithStatement( + rewritten.Sql, {"CREATE", "DROP", "ALTER", "GRANT", "REVOKE"}); + const std::string queryText = Conn_->WrapQueryForCurrentCatalog(rewritten.Sql); + NQuery::TExecuteQuerySettings execSettings; + execSettings.StatsMode(NQuery::EStatsMode::Basic); + const SQLUINTEGER queryTimeoutSec = Attributes_.GetQueryTimeoutSec(); + if (queryTimeoutSec > 0) { + execSettings.ClientTimeout(TDuration::Seconds(queryTimeoutSec)); + } + const auto txSettings = Conn_->MakeTxSettings(); + if (Conn_->GetAutocommit()) { + // TS_SNAPSHOT_RW doesn't support explicit BeginTx() - we use NoTx() instead + // DDL must use NoTx() per YDB documentation + const bool isSnapshotRw = (txSettings.GetMode() == NQuery::TTxSettings::TS_SNAPSHOT_RW); + + if (isSnapshotRw || isDdl) { + return session.ExecuteQuery( + queryText, + NQuery::TTxControl::NoTx(), + params, + execSettings).ExtractValueSync(); + } + return session.ExecuteQuery( + queryText, + NQuery::TTxControl::BeginTx(txSettings).CommitTx(), + params, + execSettings).ExtractValueSync(); + } + if (!Conn_->GetTx()) { + auto beginTxResult = session.BeginTransaction(txSettings).ExtractValueSync(); + NStatusHelpers::ThrowOnError(beginTxResult); + Conn_->SetTx(beginTxResult.GetTransaction()); + } + return session.ExecuteQuery( + queryText, + NQuery::TTxControl::Tx(*Conn_->GetTx()).CommitTx(false), + params, + execSettings).ExtractValueSync(); +} + + + +SQLRETURN TStatement::Fetch() { + if (!Cursor_) { + return SQL_NO_DATA; + } + const SQLULEN maxRows = Attributes_.GetMaxRows(); + if (maxRows > 0 && RowsFetched_ >= maxRows) { + return SQL_NO_DATA; + } + const SQLULEN rowArraySize = CurrentAppRowDesc_->GetArraySize(); + SQLUSMALLINT* const statuses = ImpRowDesc_.GetArrayStatusPtr(); + SQLULEN* const fetched = ImpRowDesc_.GetRowsProcessedPtr(); + if (fetched) { + *fetched = 0; + } + if (statuses) { + std::fill_n(statuses, rowArraySize, SQL_ROW_NOROW); + } + + SQLULEN rows = 0; + SQLRETURN result = SQL_SUCCESS; + for (; rows < rowArraySize; ++rows) { + if (maxRows > 0 && RowsFetched_ >= maxRows) { + break; + } + BindingRow_ = rows; + if (!Cursor_->Fetch()) { + break; + } + FillBoundColumns(); + ++RowsFetched_; + GetDataOffsets_.assign(Cursor_->GetColumnMeta().size(), 0); + if (fetched) { + *fetched = rows + 1; + } + if (statuses) { + statuses[rows] = LastFetchRc_ == SQL_SUCCESS_WITH_INFO + ? SQL_ROW_SUCCESS_WITH_INFO + : LastFetchRc_ == SQL_SUCCESS ? SQL_ROW_SUCCESS : SQL_ROW_ERROR; + } + if (LastFetchRc_ == SQL_ERROR) { + result = SQL_ERROR; + } else if (LastFetchRc_ == SQL_SUCCESS_WITH_INFO && result == SQL_SUCCESS) { + result = SQL_SUCCESS_WITH_INFO; + } + } + BindingRow_ = 0; + return rows == 0 && result != SQL_ERROR ? SQL_NO_DATA : result; +} + +SQLRETURN TStatement::GetData(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, + SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd) { + if (!Cursor_) { + return SQL_NO_DATA; + } + if (columnNumber < 1 || columnNumber > GetDataOffsets_.size()) { + return AddError("07009", 0, "Invalid descriptor index"); + } + const SQLRETURN rc = Cursor_->GetData( + columnNumber, targetType, targetValue, bufferLength, strLenOrInd, + &GetDataOffsets_[columnNumber - 1]); + if (const char* sqlState = ConsumeLastConvertSqlState()) { + AddError(sqlState, 0, std::strcmp(sqlState, "22003") == 0 ? "Numeric value out of range" : "Conversion error"); + } + return rc; +} + +void TStatement::FillBoundColumns() { + if (!Cursor_) { + return; + } + LastFetchRc_ = SQL_SUCCESS; + const SQLULEN bindType = CurrentAppRowDesc_->GetBindType(); + const SQLULEN offset = CurrentAppRowDesc_->GetBindOffsetPtr() + ? *CurrentAppRowDesc_->GetBindOffsetPtr() + : 0; + for (SQLSMALLINT number = 1; number <= CurrentAppRowDesc_->GetRecordCount(); ++number) { + const TDescRecord* col = CurrentAppRowDesc_->FindRecord(number); + if (!col || !col->DataPtr) { + continue; + } + const SQLULEN dataStride = bindType == SQL_BIND_BY_COLUMN + ? CTypeSize(col->Type, col->OctetLength) + : bindType; + const SQLULEN indicatorStride = bindType == SQL_BIND_BY_COLUMN + ? sizeof(SQLLEN) + : bindType; + SQLPOINTER target = OffsetPointer( + static_cast(col->DataPtr), offset, BindingRow_, dataStride); + SQLLEN* indicator = OffsetPointer( + col->IndicatorPtr, offset, BindingRow_, indicatorStride); + SQLLEN* length = OffsetPointer( + col->OctetLengthPtr, offset, BindingRow_, indicatorStride); + SQLLEN convertedLength = 0; + SQLRETURN rc = Cursor_->GetData( + static_cast(number), col->Type, target, col->OctetLength, + &convertedLength); + if (convertedLength == SQL_NULL_DATA) { + if (!indicator) { + AddError("22002", 0, "Indicator variable required but not supplied"); + rc = SQL_ERROR; + } else { + *indicator = SQL_NULL_DATA; + if (length && length != indicator) { + *length = 0; + } + } + } else { + if (length) { + *length = convertedLength; + } + if (indicator && indicator != length) { + *indicator = 0; + } + } + if (rc == SQL_SUCCESS_WITH_INFO) { + AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + if (LastFetchRc_ == SQL_SUCCESS) { + LastFetchRc_ = SQL_SUCCESS_WITH_INFO; + } + } else if (rc != SQL_SUCCESS && LastFetchRc_ == SQL_SUCCESS) { + if (const char* sqlState = ConsumeLastConvertSqlState()) { + AddError(sqlState, 0, std::strcmp(sqlState, "22003") == 0 ? "Numeric value out of range" : "Conversion error"); + } + LastFetchRc_ = rc; + } + } +} + +SQLRETURN TStatement::BindCol(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd) { + if (targetValue && columnNumber < 1) { + return AddError("07009", 0, "Invalid descriptor index"); + } + if (Cursor_) { + const size_t n = Cursor_->GetColumnMeta().size(); + if (targetValue && n > 0 && static_cast(columnNumber) > n) { + return AddError("07009", 0, "Invalid descriptor index"); + } + } + + if (!targetValue) { + CurrentAppRowDesc_->RemoveRecord(static_cast(columnNumber)); + return SQL_SUCCESS; + } + TDescRecord& record = CurrentAppRowDesc_->Record(static_cast(columnNumber)); + record.Type = targetType; + record.Length = bufferLength; + record.OctetLength = bufferLength; + record.DataPtr = targetValue; + record.IndicatorPtr = strLenOrInd; + record.OctetLengthPtr = strLenOrInd; + return SQL_SUCCESS; +} + +SQLRETURN TStatement::BindParameter(SQLUSMALLINT paramNumber, + SQLSMALLINT inputOutputType, + SQLSMALLINT valueType, + SQLSMALLINT parameterType, + SQLULEN columnSize, + SQLSMALLINT decimalDigits, + SQLPOINTER parameterValuePtr, + SQLLEN bufferLength, + SQLLEN* strLenOrIndPtr) { + + if (inputOutputType != SQL_PARAM_INPUT) { + throw TOdbcException("HYC00", 0, "Only input parameters are supported"); + } + + const bool atExec = strLenOrIndPtr + && (*strLenOrIndPtr == SQL_DATA_AT_EXEC + || *strLenOrIndPtr <= SQL_LEN_DATA_AT_EXEC_OFFSET); + + if (!parameterValuePtr && !strLenOrIndPtr) { + CurrentAppParamDesc_->RemoveRecord(static_cast(paramNumber)); + ImpParamDesc_.RemoveRecord(static_cast(paramNumber)); + return SQL_SUCCESS; + } + TDescRecord& app = CurrentAppParamDesc_->Record(static_cast(paramNumber)); + app.Type = valueType; + app.Length = bufferLength; + app.OctetLength = bufferLength; + app.DataPtr = parameterValuePtr; + app.IndicatorPtr = strLenOrIndPtr; + app.OctetLengthPtr = strLenOrIndPtr; + app.ParameterType = inputOutputType; + app.AtExec = atExec; + app.AtExecComplete = false; + app.AtExecIndicator = 0; + app.AtExecChunk.clear(); + + TDescRecord& imp = ImpParamDesc_.Record(static_cast(paramNumber)); + imp.Type = parameterType; + imp.Length = static_cast(columnSize); + imp.OctetLength = static_cast(columnSize); + imp.Precision = static_cast(columnSize); + imp.Scale = decimalDigits; + imp.Nullable = SQL_NULLABLE; + imp.ParameterType = inputOutputType; + return SQL_SUCCESS; +} + +std::vector TStatement::GetBoundParams(SQLULEN paramSet) const { + std::vector params; + const SQLULEN offset = CurrentAppParamDesc_->GetBindOffsetPtr() + ? *CurrentAppParamDesc_->GetBindOffsetPtr() + : 0; + for (SQLSMALLINT number = 1; number <= ParamCount_; ++number) { + const TDescRecord* app = CurrentAppParamDesc_->FindRecord(number); + const TDescRecord* imp = ImpParamDesc_.FindRecord(number); + if (!app || !imp) { + continue; + } + SQLLEN* lengthOrIndicator = app->IndicatorPtr == app->OctetLengthPtr + ? app->IndicatorPtr + : app->OctetLengthPtr; + TBoundParam param{ + static_cast(number), imp->ParameterType, app->Type, imp->Type, + static_cast(imp->Length), imp->Scale, app->DataPtr, app->OctetLength, + lengthOrIndicator, app->AtExec, app->AtExecComplete, app->AtExecChunk}; + param = ParamAt(param, paramSet, CurrentAppParamDesc_->GetBindType(), offset); + SQLLEN* indicator = OffsetPointer( + app->IndicatorPtr, offset, paramSet, + CurrentAppParamDesc_->GetBindType() == SQL_PARAM_BIND_BY_COLUMN + ? sizeof(SQLLEN) + : CurrentAppParamDesc_->GetBindType()); + if (indicator && *indicator == SQL_NULL_DATA) { + param.StrLenOrIndPtr = indicator; + } + params.push_back(std::move(param)); + } + return params; +} + +SQLRETURN TStatement::BuildParams(NYdb::TParams& out, SQLULEN paramSet) { + ClearErrors(); + NYdb::TParamsBuilder paramsBuilder; + for (const TBoundParam& param : GetBoundParams(paramSet)) { + const std::string paramName = "$p" + std::to_string(param.ParamNumber); + if (param.AtExec) { + if (!param.AtExecComplete) { + return AddError("HY000", 0, "Missing data-at-execution parameter value"); + } + const TDescRecord* record = CurrentAppParamDesc_->FindRecord( + static_cast(param.ParamNumber)); + SQLLEN indicator = record && record->AtExecIndicator == SQL_NULL_DATA + ? SQL_NULL_DATA + : SQL_NTS; + TBoundParam tmp = param; + tmp.ParameterValuePtr = const_cast(param.AtExecChunk.data()); + tmp.StrLenOrIndPtr = &indicator; + const SQLRETURN convRc = ConvertParam(tmp, paramsBuilder.AddParam(paramName)); + if (convRc != SQL_SUCCESS) { + return AddError("07006", 0, "Unsupported or invalid ODBC parameter type for parameter " + + std::to_string(param.ParamNumber)); + } + continue; + } + const SQLRETURN convRc = ConvertParam(param, paramsBuilder.AddParam(paramName)); + if (convRc != SQL_SUCCESS) { + return AddError( + "07006", + 0, + "Unsupported or invalid ODBC parameter type for parameter " + std::to_string(param.ParamNumber) + + " (C type " + std::to_string(static_cast(param.ValueType)) + ", SQL type " + + std::to_string(static_cast(param.ParameterType)) + ")"); + } + } + out = paramsBuilder.Build(); + return SQL_SUCCESS; +} + + +SQLRETURN TStatement::NumParams(SQLSMALLINT* paramCount) { + if (!paramCount) { + throw TOdbcException("HY000", 0, "Invalid parameter"); + } + if (!IsPrepared_) { + throw TOdbcException("HY010", 0, "Function sequence error"); + } + *paramCount = ParamCount_; + return SQL_SUCCESS; +} + +void TStatement::ResetForMetadata() { + ClearErrors(); + RowsFetched_ = 0; + RowCount_ = -1; + SetCursor(nullptr); +} + +SQLRETURN TStatement::DescribeParam(SQLUSMALLINT paramNumber, SQLSMALLINT* dataTypePtr, SQLULEN* paramSizePtr, + SQLSMALLINT* decimalDigitsPtr, SQLSMALLINT* nullablePtr) { + if (!IsPrepared_) { + throw TOdbcException("HY010", 0, "Function sequence error"); + } + if (paramNumber < 1 || paramNumber > ParamCount_) { + throw TOdbcException("07009", 0, "Invalid descriptor index"); + } + const TDescRecord* record = ImpParamDesc_.FindRecord(static_cast(paramNumber)); + const SQLSMALLINT dataType = record ? record->Type : SQL_UNKNOWN_TYPE; + const SQLULEN paramSize = record ? static_cast(record->Length) : 0; + const SQLSMALLINT decimalDigits = record ? record->Scale : 0; + const SQLSMALLINT nullable = record ? record->Nullable : SQL_NULLABLE_UNKNOWN; + if (dataTypePtr) { + *dataTypePtr = dataType; + } + if (paramSizePtr) { + *paramSizePtr = paramSize; + } + if (decimalDigitsPtr) { + *decimalDigitsPtr = decimalDigits; + } + if (nullablePtr) { + *nullablePtr = nullable; + } + return SQL_SUCCESS; +} + +SQLRETURN TStatement::ParamData(SQLPOINTER* valuePtr) { + if (!valuePtr) { + throw TOdbcException("HY009", 0, "Invalid use of null pointer"); + } + if (!InAtExec_) { + return SQL_NO_DATA; + } + if (NeedDataParam_ != 0 && NeedDataTokenDelivered_) { + if (TDescRecord* record = CurrentAppParamDesc_->FindRecord( + static_cast(NeedDataParam_))) { + record->AtExecComplete = true; + } + NeedDataParam_ = 0; + NeedDataTokenDelivered_ = false; + } + const SQLUSMALLINT next = FindNextNeedDataParam(); + if (next != 0) { + NeedDataParam_ = next; + NeedDataTokenDelivered_ = true; + *valuePtr = CurrentAppParamDesc_->FindRecord(static_cast(next))->DataPtr; + return SQL_NEED_DATA; + } + InAtExec_ = false; + NeedDataParam_ = 0; + return ExecuteInternal(); +} + +SQLRETURN TStatement::PutData(SQLPOINTER data, SQLLEN strLenOrInd) { + if (!InAtExec_ || NeedDataParam_ == 0) { + throw TOdbcException("HY010", 0, "Function sequence error"); + } + TDescRecord* param = CurrentAppParamDesc_->FindRecord( + static_cast(NeedDataParam_)); + if (!param || !NeedDataTokenDelivered_) { + throw TOdbcException("HY010", 0, "Function sequence error"); + } + SQLLEN chunkLen = strLenOrInd; + if (chunkLen == SQL_NULL_DATA) { + param->AtExecIndicator = SQL_NULL_DATA; + return SQL_SUCCESS; + } + if (chunkLen == SQL_DEFAULT_PARAM) { + throw TOdbcException("07S01", 0, "Default parameters are not supported"); + } + if (chunkLen == SQL_NTS) { + if (!data) { + throw TOdbcException("HY009", 0, "Invalid use of null pointer"); + } + chunkLen = static_cast(std::strlen(static_cast(data))); + } + if (chunkLen < 0) { + throw TOdbcException("HY090", 0, "Invalid string or buffer length"); + } + if (chunkLen > 0) { + if (!data) { + throw TOdbcException("HY009", 0, "Invalid use of null pointer"); + } + param->AtExecChunk.append(static_cast(data), static_cast(chunkLen)); + } + return SQL_SUCCESS; +} + +SQLRETURN TStatement::Cancel() { + if (!Cursor_ && !InAtExec_) { + return SQL_SUCCESS; + } + SetCursor(nullptr); + InAtExec_ = false; + NeedDataParam_ = 0; + NeedDataTokenDelivered_ = false; + for (SQLSMALLINT i = 1; i <= CurrentAppParamDesc_->GetRecordCount(); ++i) { + if (TDescRecord* param = CurrentAppParamDesc_->FindRecord(i)) { + param->AtExecComplete = false; + param->AtExecIndicator = 0; + param->AtExecChunk.clear(); + } + } + RowsFetched_ = 0; + return SQL_SUCCESS; +} + +SQLRETURN TStatement::SetCursorName(const std::string& name) { + CursorName_ = name; + return SQL_SUCCESS; +} + +SQLRETURN TStatement::GetCursorName(SQLCHAR* name, SQLSMALLINT bufferLength, SQLSMALLINT* nameLengthPtr) { + return Diag::WriteOdbcString(*this, CursorName_, name, bufferLength, nameLengthPtr); +} + + +SQLRETURN TStatement::Close(bool force) { + if (!force && !Cursor_) { + throw TOdbcException("24000", 0, "Invalid handle"); + } + + SetCursor(nullptr); + RowsFetched_ = 0; + ClearErrors(); + return SQL_SUCCESS; +} + +void TStatement::UnbindColumns() { + CurrentAppRowDesc_->ClearRecords(); +} + +void TStatement::ResetParams() { + CurrentAppParamDesc_->ClearRecords(); + ImpParamDesc_.ClearRecords(); +} + +SQLRETURN TStatement::RowCount(SQLLEN* rowCount) { + if (!rowCount) { + throw TOdbcException("HY000", 0, "Invalid parameter"); + } + + *rowCount = RowCount_; + return SQL_SUCCESS; +} + +SQLRETURN TStatement::NumResultCols(SQLSMALLINT* colCount) { + if (!colCount) { + throw TOdbcException("HY000", 0, "Invalid parameter"); + } + if (!Cursor_) { + *colCount = 0; + return SQL_SUCCESS; + } + *colCount = static_cast(Cursor_->GetColumnMeta().size()); + return SQL_SUCCESS; +} + +const std::vector& TStatement::GetColumnMeta() const { + static const std::vector EmptyColumns; + return Cursor_ ? Cursor_->GetColumnMeta() : EmptyColumns; +} + +void TStatement::SetCursor(std::unique_ptr cursor) { + Cursor_ = std::move(cursor); + GetDataOffsets_.clear(); + ImpRowDesc_.ClearRecords(); + if (!Cursor_) { + return; + } + SQLSMALLINT number = 0; + for (const TColumnMeta& column : Cursor_->GetColumnMeta()) { + TDescRecord& record = ImpRowDesc_.Record(++number); + record.Name = column.Name; + record.Type = column.SqlType; + record.Length = static_cast(column.Size); + record.OctetLength = static_cast(column.Size); + record.Precision = static_cast(column.Size); + record.Scale = column.DecimalDigits; + record.Nullable = column.Nullable; + } +} + +SQLRETURN TStatement::SetStmtAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER stringLength) { + if (attr == SQL_ATTR_APP_ROW_DESC || attr == SQL_ATTR_APP_PARAM_DESC) { + TDescriptor* desc = value ? TDescriptor::FromHandle(value) : nullptr; + if (desc && (desc->GetDescType() != EDescType::Explicit + || desc->GetConnection() != Conn_)) { + return AddError("HY024", 0, "Descriptor belongs to another connection"); + } + TDescriptor*& current = attr == SQL_ATTR_APP_ROW_DESC + ? CurrentAppRowDesc_ + : CurrentAppParamDesc_; + TDescriptor* const automatic = attr == SQL_ATTR_APP_ROW_DESC + ? &AppRowDesc_ + : &AppParamDesc_; + TDescriptor* const next = desc ? desc : automatic; + if (current != next) { + TDescriptor* const previous = current; + current = next; + current->Attach(this); + if (CurrentAppRowDesc_ != previous && CurrentAppParamDesc_ != previous) { + previous->Detach(this); + } + } + return SQL_SUCCESS; + } + const SQLULEN integer = ReadIntegerAttr(value); + switch (attr) { + case SQL_ATTR_PARAM_BIND_TYPE: CurrentAppParamDesc_->SetBindType(integer); return SQL_SUCCESS; + case SQL_ATTR_PARAMSET_SIZE: + if (integer == 0) return Diag::AddInvalidAttrValue(*this, "SQL_ATTR_PARAMSET_SIZE"); + CurrentAppParamDesc_->SetArraySize(integer); return SQL_SUCCESS; + case SQL_ATTR_PARAM_BIND_OFFSET_PTR: + CurrentAppParamDesc_->SetBindOffsetPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_PARAM_OPERATION_PTR: + CurrentAppParamDesc_->SetArrayStatusPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_PARAM_STATUS_PTR: + ImpParamDesc_.SetArrayStatusPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_PARAMS_PROCESSED_PTR: + ImpParamDesc_.SetRowsProcessedPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_ROW_BIND_TYPE: CurrentAppRowDesc_->SetBindType(integer); return SQL_SUCCESS; + case SQL_ATTR_ROW_ARRAY_SIZE: + if (integer == 0) return Diag::AddInvalidAttrValue(*this, "SQL_ATTR_ROW_ARRAY_SIZE"); + CurrentAppRowDesc_->SetArraySize(integer); return SQL_SUCCESS; + case SQL_ATTR_ROW_BIND_OFFSET_PTR: + CurrentAppRowDesc_->SetBindOffsetPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_ROW_STATUS_PTR: + ImpRowDesc_.SetArrayStatusPtr(static_cast(value)); return SQL_SUCCESS; + case SQL_ATTR_ROWS_FETCHED_PTR: + ImpRowDesc_.SetRowsProcessedPtr(static_cast(value)); return SQL_SUCCESS; + default: break; + } + return Attributes_.SetStmtAttr(attr, value, stringLength, *this); +} + +SQLRETURN TStatement::GetStmtAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr) { + if (!value) { + return AddError("HY009", 0, "Invalid use of null pointer"); + } + switch (attr) { + case SQL_ATTR_APP_ROW_DESC: + *reinterpret_cast(value) = CurrentAppRowDesc_; + return SQL_SUCCESS; + case SQL_ATTR_APP_PARAM_DESC: + *reinterpret_cast(value) = CurrentAppParamDesc_; + return SQL_SUCCESS; + case SQL_ATTR_IMP_ROW_DESC: + *reinterpret_cast(value) = &ImpRowDesc_; + return SQL_SUCCESS; + case SQL_ATTR_IMP_PARAM_DESC: + *reinterpret_cast(value) = &ImpParamDesc_; + return SQL_SUCCESS; + case SQL_ATTR_PARAM_BIND_TYPE: + *static_cast(value) = CurrentAppParamDesc_->GetBindType(); return SQL_SUCCESS; + case SQL_ATTR_PARAMSET_SIZE: + *static_cast(value) = CurrentAppParamDesc_->GetArraySize(); return SQL_SUCCESS; + case SQL_ATTR_PARAM_BIND_OFFSET_PTR: + *static_cast(value) = CurrentAppParamDesc_->GetBindOffsetPtr(); return SQL_SUCCESS; + case SQL_ATTR_PARAM_OPERATION_PTR: + *static_cast(value) = CurrentAppParamDesc_->GetArrayStatusPtr(); return SQL_SUCCESS; + case SQL_ATTR_PARAM_STATUS_PTR: + *static_cast(value) = ImpParamDesc_.GetArrayStatusPtr(); return SQL_SUCCESS; + case SQL_ATTR_PARAMS_PROCESSED_PTR: + *static_cast(value) = ImpParamDesc_.GetRowsProcessedPtr(); return SQL_SUCCESS; + case SQL_ATTR_ROW_BIND_TYPE: + *static_cast(value) = CurrentAppRowDesc_->GetBindType(); return SQL_SUCCESS; + case SQL_ATTR_ROW_ARRAY_SIZE: + *static_cast(value) = CurrentAppRowDesc_->GetArraySize(); return SQL_SUCCESS; + case SQL_ATTR_ROW_BIND_OFFSET_PTR: + *static_cast(value) = CurrentAppRowDesc_->GetBindOffsetPtr(); return SQL_SUCCESS; + case SQL_ATTR_ROW_STATUS_PTR: + *static_cast(value) = ImpRowDesc_.GetArrayStatusPtr(); return SQL_SUCCESS; + case SQL_ATTR_ROWS_FETCHED_PTR: + *static_cast(value) = ImpRowDesc_.GetRowsProcessedPtr(); return SQL_SUCCESS; + default: + break; + } + return Attributes_.GetStmtAttr(attr, value, bufferLength, stringLengthPtr, *this); +} + +SQLRETURN TStatement::GetDiagField( + SQLSMALLINT recNumber, + SQLSMALLINT diagIdentifier, + SQLPOINTER diagInfoPtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) { + if (diagIdentifier == SQL_DIAG_ROW_COUNT) { + return RowCount(static_cast(diagInfoPtr)); + } + return TErrorManager::GetDiagField(recNumber, diagIdentifier, diagInfoPtr, bufferLength, stringLengthPtr); +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/statement.h b/odbc/src/statement.h new file mode 100644 index 00000000000..643ca0bd0d8 --- /dev/null +++ b/odbc/src/statement.h @@ -0,0 +1,138 @@ +#pragma once + +#include "connection.h" +#include "statement_attr.h" +#include "descriptor.h" +#include "utils/error_manager.h" +#include "utils/bindings.h" +#include "utils/cursor.h" + +#include + +#include +#include + +#include +#include +#include +#include + + +namespace NYdb::NOdbc { + +class TStatement : public TErrorManager { + friend class TDescriptor; +public: + TStatement(TConnection* conn); + ~TStatement(); + + SQLRETURN Prepare(const std::string& statementText); + SQLRETURN Execute(); + SQLRETURN ExecuteInternal(); + + SQLRETURN Fetch(); + SQLRETURN GetData(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, + SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd); + + SQLRETURN Close(bool force = false); + void UnbindColumns(); + void ResetParams(); + + SQLRETURN BindCol(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd); + SQLRETURN BindParameter(SQLUSMALLINT paramNumber, SQLSMALLINT inputOutputType, SQLSMALLINT valueType, SQLSMALLINT parameterType, SQLULEN columnSize, SQLSMALLINT decimalDigits, SQLPOINTER parameterValuePtr, SQLLEN bufferLength, SQLLEN* strLenOrIndPtr); + + SQLRETURN Columns(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + const std::string& columnName); + + SQLRETURN Tables(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + const std::string& tableType); + + SQLRETURN GetTypeInfo(SQLSMALLINT dataType); + SQLRETURN Statistics(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + SQLUSMALLINT unique, + SQLUSMALLINT accuracy); + SQLRETURN SpecialColumns(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + SQLUSMALLINT identifierType, + SQLUSMALLINT scope); + SQLRETURN PrimaryKeys(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName); + SQLRETURN ForeignKeys(const std::string& pkCatalogName, + const std::string& pkSchemaName, + const std::string& pkTableName, + const std::string& fkCatalogName, + const std::string& fkSchemaName, + const std::string& fkTableName); + SQLRETURN NumParams(SQLSMALLINT* paramCount); + SQLRETURN DescribeParam(SQLUSMALLINT paramNumber, SQLSMALLINT* dataTypePtr, SQLULEN* paramSizePtr, + SQLSMALLINT* decimalDigitsPtr, SQLSMALLINT* nullablePtr); + SQLRETURN ParamData(SQLPOINTER* valuePtr); + SQLRETURN PutData(SQLPOINTER data, SQLLEN strLenOrInd); + SQLRETURN Cancel(); + SQLRETURN SetCursorName(const std::string& name); + SQLRETURN GetCursorName(SQLCHAR* name, SQLSMALLINT bufferLength, SQLSMALLINT* nameLengthPtr); + + void DetachDescriptor(TDescriptor* desc); + + SQLRETURN RowCount(SQLLEN* rowCount); + SQLRETURN NumResultCols(SQLSMALLINT* colCount); + const std::vector& GetColumnMeta() const; + SQLRETURN SetStmtAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER stringLength); + SQLRETURN GetStmtAttr(SQLINTEGER attr, SQLPOINTER value, SQLINTEGER bufferLength, SQLINTEGER* stringLengthPtr); + + SQLRETURN GetDiagField(SQLSMALLINT recNumber, SQLSMALLINT diagIdentifier, SQLPOINTER diagInfoPtr, SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr) override; + +private: + TConnection* Conn_; + std::unique_ptr Cursor_; + std::string PreparedQuery_; + bool IsPrepared_ = false; + SQLSMALLINT ParamCount_ = 0; + + SQLULEN RowsFetched_ = 0; + SQLLEN RowCount_ = -1; + TStatementAttributes Attributes_; + std::string CursorName_; + TDescriptor AppRowDesc_; + TDescriptor AppParamDesc_; + TDescriptor ImpRowDesc_; + TDescriptor ImpParamDesc_; + TDescriptor* CurrentAppRowDesc_; + TDescriptor* CurrentAppParamDesc_; + SQLUSMALLINT NeedDataParam_ = 0; + bool InAtExec_ = false; + bool NeedDataTokenDelivered_ = false; + SQLRETURN LastFetchRc_ = SQL_SUCCESS; + SQLULEN BindingRow_ = 0; + std::vector GetDataOffsets_; + + SQLRETURN BuildParams(NYdb::TParams& out, SQLULEN paramSet); + SQLRETURN ExecuteParamSet(SQLULEN paramSet, std::optional& affectedRows); + void FillBoundColumns(); + std::vector GetBoundParams(SQLULEN paramSet) const; + void SetCursor(std::unique_ptr cursor); + + void ResetForMetadata(); + + SQLUSMALLINT FindNextNeedDataParam() const; + std::string GetTraversalRoot(const std::string& pattern) const; + + NQuery::TExecuteQueryResult ExecuteQuery(NQuery::TSession& session, const NYdb::TParams& params); + + NYdb::NRetry::TRetryOperationSettings MakeAutocommitRetrySettings(); + std::vector GetPatternEntries(const std::string& pattern); + SQLRETURN VisitEntry(const std::string& path, const std::string& pattern, std::vector& resultEntries); + bool IsPatternMatch(const std::string& path, const std::string& pattern); + std::optional GetTableType(NScheme::ESchemeEntryType type); +}; + +} // namespace NYdb::NOdbc diff --git a/odbc/src/statement_attr.cpp b/odbc/src/statement_attr.cpp new file mode 100644 index 00000000000..1f15001444d --- /dev/null +++ b/odbc/src/statement_attr.cpp @@ -0,0 +1,113 @@ +#include "statement_attr.h" + +#include "utils/attr.h" +#include "utils/diag.h" + +#include + +namespace NYdb { +namespace NOdbc { + +SQLRETURN TStatementAttributes::SetStmtAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER /*stringLength*/, + TErrorManager& errors) { + switch (attr) { + case SQL_ATTR_QUERY_TIMEOUT: { + const SQLINTEGER timeout = ReadIntegerAttr(value); + if (timeout < 0) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_QUERY_TIMEOUT"); + } + QueryTimeoutSec_ = static_cast(timeout); + return SQL_SUCCESS; + } + case SQL_ATTR_MAX_ROWS: { + const SQLLEN maxRows = ReadIntegerAttr(value); + if (maxRows < 0) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_MAX_ROWS"); + } + MaxRows_ = static_cast(maxRows); + return SQL_SUCCESS; + } + case SQL_ATTR_NOSCAN: { + const auto mode = ReadIntegerAttrIfIn(value, {SQL_NOSCAN_OFF, SQL_NOSCAN_ON}); + if (!mode) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_NOSCAN"); + } + NoScan_ = *mode; + return SQL_SUCCESS; + } + case SQL_ATTR_METADATA_ID: { + const auto mode = ReadIntegerAttrIfIn(value, {SQL_FALSE, SQL_TRUE}); + if (!mode) { + return Diag::AddInvalidAttrValue(errors, "SQL_ATTR_METADATA_ID"); + } + MetadataId_ = *mode; + return SQL_SUCCESS; + } + case SQL_ATTR_CURSOR_TYPE: { + const SQLULEN cursorType = ReadIntegerAttr(value); + if (cursorType != SQL_CURSOR_FORWARD_ONLY) { + return errors.AddError( + "01S02", 0, "Only SQL_CURSOR_FORWARD_ONLY is supported", SQL_SUCCESS_WITH_INFO); + } + CursorType_ = cursorType; + return SQL_SUCCESS; + } + default: + return Diag::AddNotImplemented(errors); + } +} + +SQLRETURN TStatementAttributes::GetStmtAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER /*bufferLength*/, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const { + if (!value) { + return Diag::AddNullPointer(errors); + } + if (stringLengthPtr) { + *stringLengthPtr = 0; + } + switch (attr) { + case SQL_ATTR_QUERY_TIMEOUT: + *reinterpret_cast(value) = QueryTimeoutSec_; + return SQL_SUCCESS; + case SQL_ATTR_MAX_ROWS: + *reinterpret_cast(value) = MaxRows_; + return SQL_SUCCESS; + case SQL_ATTR_NOSCAN: + *reinterpret_cast(value) = NoScan_; + return SQL_SUCCESS; + case SQL_ATTR_METADATA_ID: + *reinterpret_cast(value) = MetadataId_; + return SQL_SUCCESS; + case SQL_ATTR_CURSOR_TYPE: + *reinterpret_cast(value) = CursorType_; + return SQL_SUCCESS; + default: + return Diag::AddNotImplemented(errors); + } +} + +SQLUINTEGER TStatementAttributes::GetQueryTimeoutSec() const noexcept{ + return QueryTimeoutSec_; +} + +SQLULEN TStatementAttributes::GetMaxRows() const noexcept { + return MaxRows_; +} + +SQLULEN TStatementAttributes::GetNoScanMode() const noexcept { + return NoScan_; +} + +SQLULEN TStatementAttributes::GetMetadataId() const noexcept { + return MetadataId_; +} + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/statement_attr.h b/odbc/src/statement_attr.h new file mode 100644 index 00000000000..8b92ba10ed1 --- /dev/null +++ b/odbc/src/statement_attr.h @@ -0,0 +1,40 @@ +#pragma once + +#include "utils/error_manager.h" + +#include +#include + +namespace NYdb { +namespace NOdbc { + +class TStatementAttributes { +public: + SQLRETURN SetStmtAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER stringLength, + TErrorManager& errors); + + SQLRETURN GetStmtAttr( + SQLINTEGER attr, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) const; + + SQLUINTEGER GetQueryTimeoutSec() const noexcept; + SQLULEN GetMaxRows() const noexcept; + SQLULEN GetNoScanMode() const noexcept; + SQLULEN GetMetadataId() const noexcept; + +private: + SQLUINTEGER QueryTimeoutSec_ = 0; + SQLULEN MaxRows_ = 0; + SQLULEN NoScan_ = SQL_NOSCAN_OFF; + SQLULEN MetadataId_ = SQL_FALSE; + SQLULEN CursorType_ = SQL_CURSOR_FORWARD_ONLY; +}; + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/statement_metadata.cpp b/odbc/src/statement_metadata.cpp new file mode 100644 index 00000000000..4acc63e14b5 --- /dev/null +++ b/odbc/src/statement_metadata.cpp @@ -0,0 +1,542 @@ +#include "statement.h" + +#include "utils/types.h" +#include "utils/sql_like.h" +#include "utils/type_info_rows.h" +#include "utils/cursor.h" +#include "utils/util.h" + +#include + +#include +#include +#include +#include + +namespace NYdb { +namespace NOdbc { + +namespace { + +bool MatchesTableTypeFilter(const std::string& filter, const std::string& entryType) { + if (filter.empty()) { + return true; + } + size_t start = 0; + while (start <= filter.size()) { + const size_t comma = filter.find(',', start); + std::string token = filter.substr(start, comma == std::string::npos ? std::string::npos : comma - start); + while (!token.empty() && std::isspace(static_cast(token.front()))) { + token.erase(token.begin()); + } + while (!token.empty() && std::isspace(static_cast(token.back()))) { + token.pop_back(); + } + if (token.size() >= 2 && token.front() == '\'' && token.back() == '\'') { + token = token.substr(1, token.size() - 2); + } + if (!token.empty() && token.size() == entryType.size() && + StartsWithPrefix(entryType.c_str(), entryType.size(), token.c_str(), token.size())) { + return true; + } + if (comma == std::string::npos) { + break; + } + start = comma + 1; + } + return false; +} + +namespace NColumnsRow { +constexpr int kTableCat = 0; +constexpr int kTableSchem = 1; +constexpr int kTableName = 2; +constexpr int kColumnName = 3; +constexpr int kDataType = 4; +constexpr int kTypeName = 5; +constexpr int kColumnSize = 6; +constexpr int kBufferLength = 7; +constexpr int kDecimalDigits = 8; +constexpr int kNumPrecRadix = 9; +constexpr int kNullable = 10; +constexpr int kRemarks = 11; +constexpr int kColumnDef = 12; +constexpr int kSqlDataType = 13; +constexpr int kSqlDatetimeSub = 14; +constexpr int kCharOctetLength = 15; +constexpr int kOrdinalPosition = 16; +constexpr int kIsNullable = 17; +} // namespace NColumnsRow + +} // namespace + +SQLRETURN TStatement::Columns(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + const std::string& columnName) { + ResetForMetadata(); + + std::vector columns = { + {"TABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"COLUMN_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"DATA_TYPE", SQL_INTEGER, 0, SQL_NO_NULLS}, + {"TYPE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"COLUMN_SIZE", SQL_INTEGER, 0, SQL_NULLABLE}, + {"BUFFER_LENGTH", SQL_INTEGER, 0, SQL_NULLABLE}, + {"DECIMAL_DIGITS", SQL_INTEGER, 0, SQL_NULLABLE}, + {"NUM_PREC_RADIX", SQL_INTEGER, 0, SQL_NULLABLE}, + {"NULLABLE", SQL_INTEGER, 0, SQL_NO_NULLS}, + {"REMARKS", SQL_VARCHAR, 762, SQL_NULLABLE}, + {"COLUMN_DEF", SQL_VARCHAR, 254, SQL_NULLABLE}, + {"SQL_DATA_TYPE", SQL_INTEGER, 0, SQL_NO_NULLS}, + {"SQL_DATETIME_SUB", SQL_INTEGER, 0, SQL_NULLABLE}, + {"CHAR_OCTET_LENGTH", SQL_INTEGER, 0, SQL_NULLABLE}, + {"ORDINAL_POSITION", SQL_INTEGER, 0, SQL_NO_NULLS}, + {"IS_NULLABLE", SQL_VARCHAR, 254, SQL_NO_NULLS} + }; + + auto entries = GetPatternEntries(tableName); + + TTable table; + table.reserve(entries.size()); + + if (entries.empty()) { + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; + } + + for (const auto& entry : entries) { + if (entry.Type != NScheme::ESchemeEntryType::Table && + entry.Type != NScheme::ESchemeEntryType::ColumnTable) { + continue; + } + + auto tableClient = Conn_->GetTableClient(); + if (!tableClient) { + throw TOdbcException("HY000", 0, "No client connection"); + } + + auto status = tableClient->RetryOperationSync([this, path = entry.Name, &table, &columnName](NTable::TSession session) -> TStatus { + auto result = session.DescribeTable(path).ExtractValueSync(); + NStatusHelpers::ThrowOnError(result); + + auto columns = result.GetTableDescription().GetTableColumns(); + + auto columnMatches = [&](const NTable::TTableColumn& column) { + if (columnName.empty()) { + return true; + } + if (Attributes_.GetMetadataId() == SQL_TRUE) { + return column.Name == columnName; + } + return SqlLikeMatch(column.Name, columnName); + }; + + bool foundColumn = false; + for (size_t columnIndex = 0; columnIndex < columns.size(); ++columnIndex) { + const auto& column = columns[columnIndex]; + if (!columnMatches(column)) { + continue; + } + foundColumn = true; + + const auto sqlType = GetTypeId(column.Type); + const auto colSize = GetColumnSize(sqlType); + const auto decDigits = GetDecimalDigits(column.Type); + const auto radix = GetRadix(column.Type); + const std::optional colSizeOpt = colSize > 0 ? std::optional(static_cast(colSize)) : std::nullopt; + + table.push_back({ + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().Utf8(path).Build(), + TValueBuilder().Utf8(column.Name).Build(), + TValueBuilder().Int16(sqlType).Build(), + TValueBuilder().Utf8(column.Type.ToString()).Build(), + TValueBuilder().OptionalInt32(colSizeOpt).Build(), + TValueBuilder().OptionalInt32(colSizeOpt).Build(), + TValueBuilder().OptionalInt16(decDigits).Build(), + TValueBuilder().OptionalInt16(radix).Build(), + TValueBuilder().Int16(column.NotNull && *column.NotNull ? SQL_NO_NULLS : SQL_NULLABLE).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().Int16(sqlType).Build(), + TValueBuilder().OptionalInt16(std::nullopt).Build(), + TValueBuilder().OptionalInt32(colSizeOpt).Build(), + TValueBuilder().OptionalInt32(columnIndex + 1).Build(), + TValueBuilder().Utf8(column.NotNull && *column.NotNull ? "NO" : "YES").Build(), + }); + } + if (!foundColumn && !columnName.empty()) { + return TStatus(EStatus::SUCCESS, {}); + } + return TStatus(EStatus::SUCCESS, {}); + }); + + NStatusHelpers::ThrowOnError(status); + } + + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::Tables(const std::string& catalogName, + const std::string& schemaName, + const std::string& tableName, + const std::string& tableType) { + ResetForMetadata(); + + std::vector columns = { + {"TABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"TABLE_TYPE", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"REMARKS", SQL_VARCHAR, 254, SQL_NULLABLE} + }; + + auto entries = GetPatternEntries(tableName); + + TTable table; + table.reserve(entries.size()); + + for (const auto& entry : entries) { + const auto entryType = GetTableType(entry.Type); + if (!entryType || !MatchesTableTypeFilter(tableType, *entryType)) { + continue; + } + + table.push_back({ + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().Utf8(entry.Name).Build(), + TValueBuilder().Utf8(*entryType).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + }); + } + + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::GetTypeInfo(SQLSMALLINT dataType) { + ResetForMetadata(); + + static const std::vector columns = { + {"TYPE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"DATA_TYPE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"COLUMN_SIZE", SQL_INTEGER, 0, SQL_NULLABLE}, + {"LITERAL_PREFIX", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"LITERAL_SUFFIX", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"CREATE_PARAMS", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"NULLABLE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"CASE_SENSITIVE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"SEARCHABLE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"UNSIGNED_ATTRIBUTE", SQL_CHAR, 1, SQL_NULLABLE}, + {"FIXED_PREC_SCALE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"AUTO_UNIQUE_VALUE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"LOCAL_TYPE_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"MINIMUM_SCALE", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"MAXIMUM_SCALE", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"SQL_DATA_TYPE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"SQL_DATETIME_SUB", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"NUM_PREC_RADIX", SQL_INTEGER, 0, SQL_NULLABLE}, + {"INTERVAL_PRECISION", SQL_SMALLINT, 0, SQL_NULLABLE}, + }; + + SetCursor(CreateVirtualCursor(columns, BuildTypeInfoRows(dataType))); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::Statistics(const std::string& /*catalogName*/, + const std::string& /*schemaName*/, + const std::string& /*tableName*/, + SQLUSMALLINT /*unique*/, + SQLUSMALLINT /*accuracy*/) { + ResetForMetadata(); + + static const std::vector columns = { + {"TABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"NON_UNIQUE", SQL_CHAR, 1, SQL_NO_NULLS}, + {"INDEX_QUALIFIER", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"INDEX_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TYPE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"ORDINAL_POSITION", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"COLUMN_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"ASC_OR_DESC", SQL_CHAR, 1, SQL_NULLABLE}, + {"CARDINALITY", SQL_INTEGER, 0, SQL_NULLABLE}, + {"PAGES", SQL_INTEGER, 0, SQL_NULLABLE}, + {"FILTER_CONDITION", SQL_VARCHAR, 128, SQL_NULLABLE}, + }; + + SetCursor(CreateVirtualCursor(columns, TTable{})); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::SpecialColumns(const std::string& /*catalogName*/, + const std::string& /*schemaName*/, + const std::string& tableName, + SQLUSMALLINT identifierType, + SQLUSMALLINT /*scope*/) { + if (identifierType != SQL_BEST_ROWID) { + return AddError("HYC00", 0, "Optional feature not implemented"); + } + + ResetForMetadata(); + + std::vector columns = { + {"SCOPE", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"COLUMN_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"DATA_TYPE", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"TYPE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"COLUMN_SIZE", SQL_INTEGER, 0, SQL_NULLABLE}, + {"BUFFER_LENGTH", SQL_INTEGER, 0, SQL_NULLABLE}, + {"DECIMAL_DIGITS", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"PSEUDO_COLUMN", SQL_SMALLINT, 0, SQL_NO_NULLS}, + }; + + TTable table; + auto entries = GetPatternEntries(tableName); + if (entries.size() != 1) { + if (entries.empty()) { + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; + } + throw TOdbcException("HY000", 0, "Ambiguous table name"); + } + + auto tableClient = Conn_->GetTableClient(); + if (!tableClient) { + throw TOdbcException("HY000", 0, "No client connection"); + } + + const std::string path = entries.front().Name; + auto status = tableClient->RetryOperationSync([path, &table, &columns](NTable::TSession session) -> TStatus { + auto result = session.DescribeTable(path).ExtractValueSync(); + NStatusHelpers::ThrowOnError(result); + + const auto& pkColumns = result.GetTableDescription().GetPrimaryKeyColumns(); + const auto& tableColumns = result.GetTableDescription().GetTableColumns(); + for (const auto& pkName : pkColumns) { + const auto columnIt = std::ranges::find_if(tableColumns, + [&](const NTable::TTableColumn& column) { return column.Name == pkName; }); + if (columnIt == tableColumns.end()) { + continue; + } + const auto sqlType = GetTypeId(columnIt->Type); + const auto colSize = GetColumnSize(sqlType); + const std::optional colSizeOpt = colSize > 0 ? std::optional(static_cast(colSize)) : std::nullopt; + table.push_back({ + TValueBuilder().OptionalInt16(SQL_SCOPE_SESSION).Build(), + TValueBuilder().Utf8(pkName).Build(), + TValueBuilder().Int16(sqlType).Build(), + TValueBuilder().Utf8(columnIt->Type.ToString()).Build(), + TValueBuilder().OptionalInt32(colSizeOpt).Build(), + TValueBuilder().OptionalInt32(colSizeOpt).Build(), + TValueBuilder().OptionalInt16(GetDecimalDigits(columnIt->Type)).Build(), + TValueBuilder().Int16(SQL_PC_NOT_PSEUDO).Build(), + }); + } + return TStatus(EStatus::SUCCESS, {}); + }); + NStatusHelpers::ThrowOnError(status); + + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::PrimaryKeys(const std::string& /*catalogName*/, + const std::string& /*schemaName*/, + const std::string& tableName) { + ResetForMetadata(); + + std::vector columns = { + {"TABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"TABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"COLUMN_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"KEY_SEQ", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"PK_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + }; + + TTable table; + auto entries = GetPatternEntries(tableName); + if (entries.size() != 1) { + if (entries.empty()) { + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; + } + throw TOdbcException("HY000", 0, "Ambiguous table name"); + } + + auto tableClient = Conn_->GetTableClient(); + if (!tableClient) { + throw TOdbcException("HY000", 0, "No client connection"); + } + + const std::string path = entries.front().Name; + auto status = tableClient->RetryOperationSync([path, &table](NTable::TSession session) -> TStatus { + auto result = session.DescribeTable(path).ExtractValueSync(); + NStatusHelpers::ThrowOnError(result); + + const auto& pkColumns = result.GetTableDescription().GetPrimaryKeyColumns(); + SQLSMALLINT keySeq = 1; + for (const auto& pkName : pkColumns) { + table.push_back({ + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + TValueBuilder().Utf8(path).Build(), + TValueBuilder().Utf8(pkName).Build(), + TValueBuilder().Int16(keySeq++).Build(), + TValueBuilder().OptionalUtf8(std::nullopt).Build(), + }); + } + return TStatus(EStatus::SUCCESS, {}); + }); + NStatusHelpers::ThrowOnError(status); + + SetCursor(CreateVirtualCursor(columns, table)); + return SQL_SUCCESS; +} + +SQLRETURN TStatement::ForeignKeys(const std::string& /*pkCatalogName*/, + const std::string& /*pkSchemaName*/, + const std::string& /*pkTableName*/, + const std::string& /*fkCatalogName*/, + const std::string& /*fkSchemaName*/, + const std::string& /*fkTableName*/) { + ResetForMetadata(); + + std::vector columns = { + {"PKTABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"PKTABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"PKTABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"PKCOLUMN_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"FKTABLE_CAT", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"FKTABLE_SCHEM", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"FKTABLE_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"FKCOLUMN_NAME", SQL_VARCHAR, 128, SQL_NO_NULLS}, + {"KEY_SEQ", SQL_SMALLINT, 0, SQL_NO_NULLS}, + {"UPDATE_RULE", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"DELETE_RULE", SQL_SMALLINT, 0, SQL_NULLABLE}, + {"FK_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"PK_NAME", SQL_VARCHAR, 128, SQL_NULLABLE}, + {"DEFERRABILITY", SQL_SMALLINT, 0, SQL_NULLABLE}, + }; + + SetCursor(CreateVirtualCursor(columns, TTable{})); + return SQL_SUCCESS; +} + +std::string TStatement::GetTraversalRoot(const std::string& pattern) const { + if (pattern.empty()) { + return ""; + } + const auto hasWildcard = [](const std::string& value) { + return value.find('%') != std::string::npos || value.find('_') != std::string::npos; + }; + if (Attributes_.GetMetadataId() == SQL_TRUE && !hasWildcard(pattern)) { + const auto pos = pattern.find_last_of('/'); + return pos == std::string::npos ? "" : pattern.substr(0, pos); + } + size_t wildPos = pattern.size(); + const auto pct = pattern.find('%'); + const auto usc = pattern.find('_'); + if (pct != std::string::npos) { + wildPos = std::min(wildPos, pct); + } + if (usc != std::string::npos) { + wildPos = std::min(wildPos, usc); + } + const std::string prefix = pattern.substr(0, wildPos); + const auto pos = prefix.find_last_of('/'); + return pos == std::string::npos ? "" : prefix.substr(0, pos); +} + +std::vector TStatement::GetPatternEntries(const std::string& pattern) { + std::vector entries; + VisitEntry(GetTraversalRoot(pattern), pattern, entries); + return entries; +} + +SQLRETURN TStatement::VisitEntry(const std::string& path, const std::string& pattern, std::vector& resultEntries) { + auto schemeClient = Conn_->GetSchemeClient(); + if (!schemeClient) { + throw TOdbcException("HY000", 0, "No client connection"); + } + auto listDirectoryResult = schemeClient->ListDirectory(path + "/").ExtractValueSync(); + NStatusHelpers::ThrowOnError(listDirectoryResult); + + for (const auto& entry : listDirectoryResult.GetChildren()) { + std::string fullPath = path + "/" + entry.Name; + if (entry.Type == NScheme::ESchemeEntryType::Directory || + entry.Type == NScheme::ESchemeEntryType::SubDomain) { + VisitEntry(fullPath, pattern, resultEntries); + } else if (IsPatternMatch(fullPath, pattern)) { + NScheme::TSchemeEntry entryCopy = entry; + entryCopy.Name = fullPath; + resultEntries.push_back(entryCopy); + } + } + return SQL_SUCCESS; +} + +bool TStatement::IsPatternMatch(const std::string& path, const std::string& pattern) { + if (pattern.empty()) { + return true; + } + if (Attributes_.GetMetadataId() == SQL_TRUE) { + return path == pattern; + } + return SqlLikeMatch(path, pattern); +} + +std::optional TStatement::GetTableType(NScheme::ESchemeEntryType type) { + switch (type) { + case NScheme::ESchemeEntryType::Table: + return "TABLE"; + case NScheme::ESchemeEntryType::View: + return "VIEW"; + case NScheme::ESchemeEntryType::ColumnStore: + return "COLUMN_STORE"; + case NScheme::ESchemeEntryType::ColumnTable: + return "COLUMN_TABLE"; + case NScheme::ESchemeEntryType::Sequence: + return "SEQUENCE"; + case NScheme::ESchemeEntryType::Replication: + return "REPLICATION"; + case NScheme::ESchemeEntryType::Topic: + return "TOPIC"; + case NScheme::ESchemeEntryType::ExternalTable: + return "EXTERNAL_TABLE"; + case NScheme::ESchemeEntryType::ExternalDataSource: + return "EXTERNAL_DATA_SOURCE"; + case NScheme::ESchemeEntryType::ResourcePool: + return "RESOURCE_POOL"; + case NScheme::ESchemeEntryType::PqGroup: + return "PQ_GROUP"; + case NScheme::ESchemeEntryType::RtmrVolume: + return "RTMR_VOLUME"; + case NScheme::ESchemeEntryType::BlockStoreVolume: + return "BLOCK_STORE_VOLUME"; + case NScheme::ESchemeEntryType::CoordinationNode: + return "COORDINATION_NODE"; + case NScheme::ESchemeEntryType::Unknown: + return "UNKNOWN"; + case NScheme::ESchemeEntryType::SysView: + return "SYSTEM VIEW"; + case NScheme::ESchemeEntryType::Transfer: + return "TRANSFER"; + case NScheme::ESchemeEntryType::Directory: + case NScheme::ESchemeEntryType::SubDomain: + return std::nullopt; + default: + return std::nullopt; + } +} + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/attr.cpp b/odbc/src/utils/attr.cpp new file mode 100644 index 00000000000..1fb2a83324a --- /dev/null +++ b/odbc/src/utils/attr.cpp @@ -0,0 +1,51 @@ +#include "attr.h" +#include "diag.h" + +#include +#include + +namespace NYdb::NOdbc { + +std::string ReadAttributeString(SQLPOINTER value, SQLINTEGER stringLength) { + const char* const str = static_cast(value); + if (stringLength == SQL_NTS) { + return std::string(str); + } + if (stringLength < 0) { + return {}; + } + return std::string(str, static_cast(stringLength)); +} + +SQLRETURN WriteAttributeString( + const std::string& source, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors) { + const SQLINTEGER length = static_cast(source.size()); + if (stringLengthPtr != nullptr) { + *stringLengthPtr = length; + } + if (value == nullptr) { + return SQL_SUCCESS; + } + if (bufferLength <= 0) { + return Diag::AddInvalidBufferLength(errors); + } + + auto* dest = static_cast(value); + const size_t maxData = static_cast(bufferLength - 1); + const size_t nCopy = std::min(source.size(), maxData); + if (nCopy > 0) { + std::memcpy(dest, source.data(), nCopy); + } + dest[nCopy] = 0; + + if (length >= bufferLength) { + return Diag::AddRightTruncated(errors); + } + return SQL_SUCCESS; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/attr.h b/odbc/src/utils/attr.h new file mode 100644 index 00000000000..96695c221aa --- /dev/null +++ b/odbc/src/utils/attr.h @@ -0,0 +1,39 @@ +#pragma once + +#include "error_manager.h" + +#include +#include +#include + +#include +#include + +namespace NYdb::NOdbc { + +std::string ReadAttributeString(SQLPOINTER value, SQLINTEGER stringLength); + +SQLRETURN WriteAttributeString( + const std::string& source, + SQLPOINTER value, + SQLINTEGER bufferLength, + SQLINTEGER* stringLengthPtr, + TErrorManager& errors); + +template +T ReadIntegerAttr(SQLPOINTER value) noexcept { + return static_cast(reinterpret_cast(value)); +} + +template +std::optional ReadIntegerAttrIfIn(SQLPOINTER value, std::initializer_list allowed) noexcept { + const T token = ReadIntegerAttr(value); + for (const T allowedValue : allowed) { + if (token == allowedValue) { + return token; + } + } + return std::nullopt; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/bindings.h b/odbc/src/utils/bindings.h new file mode 100644 index 00000000000..17ce8158609 --- /dev/null +++ b/odbc/src/utils/bindings.h @@ -0,0 +1,29 @@ +#pragma once + +#include +#include + +#include + +#include + +namespace NYdb { +namespace NOdbc { + +struct TBoundParam { + SQLUSMALLINT ParamNumber; + SQLSMALLINT InputOutputType; + SQLSMALLINT ValueType; + SQLSMALLINT ParameterType; + SQLULEN ColumnSize; + SQLSMALLINT DecimalDigits; + SQLPOINTER ParameterValuePtr; + SQLLEN BufferLength; + SQLLEN* StrLenOrIndPtr; + bool AtExec = false; + bool AtExecComplete = false; + std::string AtExecChunk; +}; + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/convert.cpp b/odbc/src/utils/convert.cpp new file mode 100644 index 00000000000..46e658beda2 --- /dev/null +++ b/odbc/src/utils/convert.cpp @@ -0,0 +1,302 @@ +#include "convert.h" + +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { +namespace { + +thread_local const char* LastConvertSqlState = nullptr; + +void SetNumericOutOfRange() { + LastConvertSqlState = "22003"; +} + +bool FitsInt16(int64_t value) { + return value >= INT16_MIN && value <= INT16_MAX; +} + +bool FitsInt32(int64_t value) { + return value >= INT32_MIN && value <= INT32_MAX; +} + +std::optional GetAsInt64(TValueParser& parser, EPrimitiveType type) { + switch (type) { + case EPrimitiveType::Bool: return parser.GetBool() ? 1 : 0; + case EPrimitiveType::Int8: return parser.GetInt8(); + case EPrimitiveType::Uint8: return parser.GetUint8(); + case EPrimitiveType::Int16: return parser.GetInt16(); + case EPrimitiveType::Uint16: return parser.GetUint16(); + case EPrimitiveType::Int32: return parser.GetInt32(); + case EPrimitiveType::Uint32: return parser.GetUint32(); + case EPrimitiveType::Int64: return parser.GetInt64(); + case EPrimitiveType::Uint64: { + const uint64_t value = parser.GetUint64(); + if (value <= static_cast(INT64_MAX)) return static_cast(value); + SetNumericOutOfRange(); + return std::nullopt; + } + default: return std::nullopt; + } +} + +std::optional ReadInteger(const TBoundParam& param) { + if (!param.ParameterValuePtr) return std::nullopt; + const auto type = param.ValueType; + if (type == SQL_C_SBIGINT) return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_UBIGINT) { + const SQLUBIGINT value = *static_cast(param.ParameterValuePtr); + if (value <= static_cast(INT64_MAX)) return static_cast(value); + SetNumericOutOfRange(); + return std::nullopt; + } + if (type == SQL_C_LONG || type == SQL_C_SLONG) + return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_ULONG) + return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_SHORT || type == SQL_C_SSHORT) + return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_USHORT) + return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_TINYINT || type == SQL_C_STINYINT) + return *static_cast(param.ParameterValuePtr); + if (type == SQL_C_UTINYINT || type == SQL_C_BIT) + return *static_cast(param.ParameterValuePtr); + return std::nullopt; +} + +std::optional ReadBytes(const TBoundParam& param) { + if (!param.ParameterValuePtr) return std::nullopt; + const char* data = static_cast(param.ParameterValuePtr); + SQLLEN length = param.BufferLength; + if (param.StrLenOrIndPtr) { + length = *param.StrLenOrIndPtr; + if (length == SQL_NTS) return std::string(data); + if (length < 0) length = param.BufferLength; + } + if (length < 0) return std::nullopt; + return std::string(data, static_cast(length)); +} + +std::optional ParameterPrimitive(SQLSMALLINT sqlType) { + switch (sqlType) { + case SQL_BIGINT: return EPrimitiveType::Int64; + case SQL_INTEGER: return EPrimitiveType::Int32; + case SQL_SMALLINT: return EPrimitiveType::Int16; + case SQL_TINYINT: return EPrimitiveType::Int8; + case SQL_BIT: return EPrimitiveType::Bool; + case SQL_REAL: return EPrimitiveType::Float; + case SQL_FLOAT: + case SQL_DOUBLE: return EPrimitiveType::Double; + case SQL_CHAR: + case SQL_VARCHAR: + case SQL_LONGVARCHAR: return EPrimitiveType::Utf8; + case SQL_BINARY: + case SQL_VARBINARY: + case SQL_LONGVARBINARY: return EPrimitiveType::String; + default: return std::nullopt; + } +} + +bool IsNull(const TBoundParam& param) { + return param.StrLenOrIndPtr && *param.StrLenOrIndPtr == SQL_NULL_DATA; +} + +} // namespace + +SQLRETURN ConvertParam(const TBoundParam& param, TParamValueBuilder& builder) { + const auto primitive = ParameterPrimitive(param.ParameterType); + if (!primitive) return SQL_ERROR; + if (IsNull(param)) { + builder.EmptyOptional(TTypeBuilder().Primitive(*primitive).Build()).Build(); + return SQL_SUCCESS; + } + + if (param.ParameterType == SQL_BIGINT || param.ParameterType == SQL_INTEGER + || param.ParameterType == SQL_SMALLINT || param.ParameterType == SQL_TINYINT + || param.ParameterType == SQL_BIT) { + const auto value = ReadInteger(param); + if (!value) return SQL_ERROR; + switch (param.ParameterType) { + case SQL_BIGINT: builder.OptionalInt64(*value); break; + case SQL_INTEGER: + if (!FitsInt32(*value)) { SetNumericOutOfRange(); return SQL_ERROR; } + builder.OptionalInt32(static_cast(*value)); break; + case SQL_SMALLINT: + if (!FitsInt16(*value)) { SetNumericOutOfRange(); return SQL_ERROR; } + builder.OptionalInt16(static_cast(*value)); break; + case SQL_TINYINT: + if (*value < INT8_MIN || *value > INT8_MAX) { SetNumericOutOfRange(); return SQL_ERROR; } + builder.OptionalInt8(static_cast(*value)); break; + case SQL_BIT: + if (*value != 0 && *value != 1) { SetNumericOutOfRange(); return SQL_ERROR; } + builder.OptionalBool(*value != 0); break; + } + } else if (param.ParameterType == SQL_REAL) { + if (param.ValueType != SQL_C_FLOAT || !param.ParameterValuePtr) return SQL_ERROR; + builder.OptionalFloat(*static_cast(param.ParameterValuePtr)); + } else if (param.ParameterType == SQL_FLOAT || param.ParameterType == SQL_DOUBLE) { + if (param.ValueType != SQL_C_DOUBLE || !param.ParameterValuePtr) return SQL_ERROR; + builder.OptionalDouble(*static_cast(param.ParameterValuePtr)); + } else { + const auto bytes = ReadBytes(param); + if (!bytes) return SQL_ERROR; + if (param.ValueType == SQL_C_CHAR + && (param.ParameterType == SQL_CHAR || param.ParameterType == SQL_VARCHAR + || param.ParameterType == SQL_LONGVARCHAR)) { + builder.OptionalUtf8(*bytes); + } else if (param.ValueType == SQL_C_BINARY + && (param.ParameterType == SQL_BINARY || param.ParameterType == SQL_VARBINARY + || param.ParameterType == SQL_LONGVARBINARY)) { + builder.OptionalString(*bytes); + } else { + return SQL_ERROR; + } + } + builder.Build(); + return SQL_SUCCESS; +} + +SQLRETURN ConvertColumn(TValueParser& parser, SQLSMALLINT targetType, SQLPOINTER targetValue, + SQLLEN bufferLength, SQLLEN* strLenOrInd, SQLLEN* offset) { + LastConvertSqlState = nullptr; + if (bufferLength < 0) { + LastConvertSqlState = "HY090"; + return SQL_ERROR; + } + if (parser.IsNull()) { + if (!strLenOrInd) { + LastConvertSqlState = "22002"; + return SQL_ERROR; + } + *strLenOrInd = SQL_NULL_DATA; + return SQL_SUCCESS; + } + if (parser.GetKind() == TTypeParser::ETypeKind::Optional) { + parser.OpenOptional(); + const SQLRETURN result = ConvertColumn( + parser, targetType, targetValue, bufferLength, strLenOrInd, offset); + parser.CloseOptional(); + return result; + } + if (parser.GetKind() != TTypeParser::ETypeKind::Primitive) return SQL_ERROR; + const EPrimitiveType ydbType = parser.GetPrimitiveType(); + + if (targetType == SQL_C_SHORT || targetType == SQL_C_SSHORT + || targetType == SQL_C_LONG || targetType == SQL_C_SLONG + || targetType == SQL_C_SBIGINT || targetType == SQL_C_BIT) { + const auto raw = GetAsInt64(parser, ydbType); + if (!raw) return SQL_ERROR; + if (targetType == SQL_C_SHORT || targetType == SQL_C_SSHORT) { + if (!FitsInt16(*raw)) { SetNumericOutOfRange(); return SQL_ERROR; } + if (targetValue) *static_cast(targetValue) = static_cast(*raw); + if (strLenOrInd) *strLenOrInd = sizeof(SQLSMALLINT); + } else if (targetType == SQL_C_LONG || targetType == SQL_C_SLONG) { + if (!FitsInt32(*raw)) { SetNumericOutOfRange(); return SQL_ERROR; } + if (targetValue) *static_cast(targetValue) = static_cast(*raw); + if (strLenOrInd) *strLenOrInd = sizeof(SQLINTEGER); + } else if (targetType == SQL_C_SBIGINT) { + if (targetValue) *static_cast(targetValue) = *raw; + if (strLenOrInd) *strLenOrInd = sizeof(SQLBIGINT); + } else { + if (*raw != 0 && *raw != 1) { SetNumericOutOfRange(); return SQL_ERROR; } + if (targetValue) *static_cast(targetValue) = *raw != 0; + if (strLenOrInd) *strLenOrInd = sizeof(SQLCHAR); + } + return SQL_SUCCESS; + } + if (targetType == SQL_C_DOUBLE) { + double value; + if (ydbType == EPrimitiveType::Double) value = parser.GetDouble(); + else if (ydbType == EPrimitiveType::Float) value = parser.GetFloat(); + else return SQL_ERROR; + if (targetValue) *static_cast(targetValue) = value; + if (strLenOrInd) *strLenOrInd = sizeof(SQLDOUBLE); + return SQL_SUCCESS; + } + if (targetType != SQL_C_CHAR) return SQL_ERROR; + + std::string text; + switch (ydbType) { + case EPrimitiveType::Utf8: text = parser.GetUtf8(); break; + case EPrimitiveType::String: text = parser.GetString(); break; + case EPrimitiveType::Json: text = parser.GetJson(); break; + case EPrimitiveType::JsonDocument: text = parser.GetJsonDocument(); break; + case EPrimitiveType::Bool: text = parser.GetBool() ? "1" : "0"; break; + case EPrimitiveType::Int8: text = std::to_string(parser.GetInt8()); break; + case EPrimitiveType::Uint8: text = std::to_string(parser.GetUint8()); break; + case EPrimitiveType::Int16: text = std::to_string(parser.GetInt16()); break; + case EPrimitiveType::Uint16: text = std::to_string(parser.GetUint16()); break; + case EPrimitiveType::Int32: text = std::to_string(parser.GetInt32()); break; + case EPrimitiveType::Uint32: text = std::to_string(parser.GetUint32()); break; + case EPrimitiveType::Int64: text = std::to_string(parser.GetInt64()); break; + case EPrimitiveType::Uint64: text = std::to_string(parser.GetUint64()); break; + case EPrimitiveType::Float: text = std::to_string(parser.GetFloat()); break; + case EPrimitiveType::Double: text = std::to_string(parser.GetDouble()); break; + case EPrimitiveType::Date: { + const TString value = parser.GetDate().FormatGmTime("%Y-%m-%d"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::Date32: { + const auto days = parser.GetDate32().time_since_epoch().count(); + if (days < 0) return SQL_ERROR; + const TString value = TInstant::Days(static_cast(days)).FormatGmTime("%Y-%m-%d"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::Datetime: { + const TString value = parser.GetDatetime().FormatGmTime("%Y-%m-%d %H:%M:%S"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::Datetime64: { + const auto seconds = parser.GetDatetime64().time_since_epoch().count(); + if (seconds < 0) return SQL_ERROR; + const TString value = TInstant::Seconds(static_cast(seconds)) + .FormatGmTime("%Y-%m-%d %H:%M:%S"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::Timestamp: { + const TString value = parser.GetTimestamp().FormatGmTime("%Y-%m-%d %H:%M:%S"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::Timestamp64: { + const auto micros = parser.GetTimestamp64().time_since_epoch().count(); + if (micros < 0) return SQL_ERROR; + const TString value = TInstant::MicroSeconds(static_cast(micros)) + .FormatGmTime("%Y-%m-%d %H:%M:%S"); + text.assign(value.data(), value.size()); break; + } + case EPrimitiveType::TzDate: text = parser.GetTzDate(); break; + case EPrimitiveType::TzDatetime: text = parser.GetTzDatetime(); break; + case EPrimitiveType::TzTimestamp: text = parser.GetTzTimestamp(); break; + default: return SQL_ERROR; + } + + if (offset && *offset < 0) return SQL_NO_DATA; + const SQLLEN start = offset ? *offset : 0; + const SQLLEN remaining = static_cast(text.size()) - start; + if (targetValue && bufferLength > 0) { + const SQLLEN copied = std::min(remaining, bufferLength - 1); + std::memcpy(targetValue, text.data() + start, static_cast(copied)); + static_cast(targetValue)[copied] = 0; + if (offset) *offset = copied == remaining ? -1 : start + copied; + } + if (strLenOrInd) *strLenOrInd = remaining; + return targetValue && bufferLength > 0 && remaining >= bufferLength + ? SQL_SUCCESS_WITH_INFO + : SQL_SUCCESS; +} + +const char* ConsumeLastConvertSqlState() { + const char* result = LastConvertSqlState; + LastConvertSqlState = nullptr; + return result; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/convert.h b/odbc/src/utils/convert.h new file mode 100644 index 00000000000..70e098e2e68 --- /dev/null +++ b/odbc/src/utils/convert.h @@ -0,0 +1,19 @@ +#pragma once + +#include "bindings.h" + +#include + +#include +#include + +namespace NYdb { +namespace NOdbc { + +SQLRETURN ConvertParam(const TBoundParam& param, TParamValueBuilder& builder); +SQLRETURN ConvertColumn(TValueParser& parser, SQLSMALLINT targetType, SQLPOINTER targetValue, + SQLLEN bufferLength, SQLLEN* strLenOrInd, SQLLEN* offset = nullptr); +const char* ConsumeLastConvertSqlState(); + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/cursor.cpp b/odbc/src/utils/cursor.cpp new file mode 100644 index 00000000000..b145639b3f2 --- /dev/null +++ b/odbc/src/utils/cursor.cpp @@ -0,0 +1,94 @@ +#include "cursor.h" +#include "convert.h" +#include "types.h" + +#include + +namespace NYdb { +namespace NOdbc { + +class TExecCursor : public ICursor { +public: + explicit TExecCursor(TResultSet resultSet) + : Parser_(resultSet) { + for (const auto& col : resultSet.GetColumnsMeta()) { + const SQLSMALLINT sqlType = GetTypeId(col.Type); + Columns_.push_back({col.Name, sqlType, GetColumnSize(sqlType), IsNullable(col.Type), + GetDecimalDigits(col.Type).value_or(0)}); + } + } + + bool Fetch() override { + return Parser_.TryNextRow(); + } + + SQLRETURN GetData(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, + SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd, + SQLLEN* offset) override { + if (columnNumber < 1 || columnNumber > Parser_.ColumnsCount()) { + return SQL_ERROR; + } + return ConvertColumn( + Parser_.ColumnParser(columnNumber - 1), targetType, targetValue, bufferLength, strLenOrInd, + offset); + } + + const std::vector& GetColumnMeta() const override { + return Columns_; + } + +private: + TResultSetParser Parser_; + std::vector Columns_; +}; + +class TVirtualCursor : public ICursor { +public: + TVirtualCursor(const std::vector& columns, const TTable& table) + : Columns_(columns) + , Table_(table) + {} + + bool Fetch() override { + Cursor_++; + if (Cursor_ >= static_cast(Table_.size())) { + return false; + } + return true; + } + + SQLRETURN GetData(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, + SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd, + SQLLEN* offset) override { + if (Cursor_ >= static_cast(Table_.size())) { + return SQL_NO_DATA; + } + if (Cursor_ < 0 || columnNumber < 1 || columnNumber > Columns_.size()) { + return SQL_ERROR; + } + TValueParser parser{Table_[Cursor_][columnNumber - 1]}; + return ConvertColumn(parser, targetType, targetValue, bufferLength, strLenOrInd, offset); + } + + const std::vector& GetColumnMeta() const override { + return Columns_; + } + +private: + std::vector Columns_; + TTable Table_; + int64_t Cursor_ = -1; +}; + +std::unique_ptr CreateExecCursor(const NQuery::TExecuteQueryResult& result) { + return result.GetResultSets().empty() + ? nullptr + : std::make_unique(result.GetResultSet(0)); +} + +std::unique_ptr CreateVirtualCursor(const std::vector& columns, const TTable& table) { + return std::make_unique(columns, table); +} + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/cursor.h b/odbc/src/utils/cursor.h new file mode 100644 index 00000000000..3032eb050f9 --- /dev/null +++ b/odbc/src/utils/cursor.h @@ -0,0 +1,45 @@ +#pragma once + +#include "bindings.h" + +#include +#include + +#include + +#include +#include +#include +#include + +namespace NYdb { +namespace NOdbc { + +struct TColumnMeta { + std::string Name; + SQLSMALLINT SqlType; + SQLULEN Size; + SQLSMALLINT Nullable; + SQLSMALLINT DecimalDigits = 0; +}; + +using TTable = std::vector>; + +class ICursor { +public: + virtual ~ICursor() = default; + virtual bool Fetch() = 0; + virtual SQLRETURN GetData(SQLUSMALLINT columnNumber, SQLSMALLINT targetType, + SQLPOINTER targetValue, SQLLEN bufferLength, SQLLEN* strLenOrInd, + SQLLEN* offset = nullptr) = 0; + virtual const std::vector& GetColumnMeta() const = 0; +}; + +std::unique_ptr CreateExecCursor(const NYdb::NQuery::TExecuteQueryResult& result); + +std::unique_ptr CreateVirtualCursor( + const std::vector& columns, + const TTable& table); + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/diag.h b/odbc/src/utils/diag.h new file mode 100644 index 00000000000..3b87baae004 --- /dev/null +++ b/odbc/src/utils/diag.h @@ -0,0 +1,67 @@ +#pragma once + +#include "error_manager.h" + +#include +#include +#include +#include + +namespace NYdb::NOdbc { +namespace Diag { + + inline SQLRETURN AddNullPointer(TErrorManager& errors) { + return errors.AddError("HY009", 0, "Invalid use of null pointer"); + } + + inline SQLRETURN AddNotImplemented(TErrorManager& errors) { + return errors.AddError("HYC00", 0, "Optional feature not implemented"); + } + + inline SQLRETURN AddInvalidAttrValue(TErrorManager& errors, std::string_view attrName) { + return errors.AddError("HY024", 0, "Invalid " + std::string(attrName) + " value"); + } + + inline SQLRETURN AddInvalidBufferLength(TErrorManager& errors) { + return errors.AddError("HY090", 0, "Invalid string or buffer length"); + } + + inline SQLRETURN AddRightTruncated(TErrorManager& errors) { + return errors.AddError("01004", 0, "String data, right truncated", SQL_SUCCESS_WITH_INFO); + } + + inline SQLRETURN WriteOdbcString( + TErrorManager& errors, + std::string_view value, + SQLPOINTER outPtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* lengthPtr) { + if (!outPtr) { + return errors.AddError("HY009", 0, "Invalid use of null pointer"); + } + if (bufferLength < 0) { + return errors.AddError("HY090", 0, "Invalid string or buffer length"); + } + const SQLLEN fullLen = static_cast(value.size()); + const SQLSMALLINT reportedLen = static_cast(std::min(fullLen, 32767)); + if (lengthPtr) { + *lengthPtr = reportedLen; + } + if (bufferLength == 0) { + return fullLen == 0 ? SQL_SUCCESS : AddRightTruncated(errors); + } + auto* out = reinterpret_cast(outPtr); + const SQLSMALLINT copyLen = static_cast(std::min(fullLen, static_cast(bufferLength - 1))); + if (copyLen > 0) { + std::memcpy(out, value.data(), static_cast(copyLen)); + } + out[copyLen] = '\0'; + if (copyLen < fullLen) { + return AddRightTruncated(errors); + } + return SQL_SUCCESS; + } + +} // namespace Diag + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/error_manager.cpp b/odbc/src/utils/error_manager.cpp new file mode 100644 index 00000000000..86eedb05b51 --- /dev/null +++ b/odbc/src/utils/error_manager.cpp @@ -0,0 +1,224 @@ +#include "error_manager.h" + +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +namespace { + struct OdbcErrorMapping { + const char* sqlState; + const char* description; + SQLRETURN returnCode; + }; + + const std::unordered_map ERROR_MAPPINGS = { + {EStatus::SUCCESS, {"00000", "Success", SQL_SUCCESS}}, + {EStatus::BAD_REQUEST, {"42000", "Syntax error or access rule violation", SQL_ERROR}}, + {EStatus::UNAUTHORIZED, {"28000", "Invalid authorization specification", SQL_ERROR}}, + {EStatus::INTERNAL_ERROR, {"HY000", "General error", SQL_ERROR}}, + {EStatus::ABORTED, {"40001", "Serialization failure", SQL_ERROR}}, + {EStatus::UNAVAILABLE, {"08S01", "Communication link failure", SQL_ERROR}}, + {EStatus::OVERLOADED, {"HY000", "General error - server overloaded", SQL_ERROR}}, + {EStatus::SCHEME_ERROR, {"42S02", "Base table or view not found", SQL_ERROR}}, + {EStatus::GENERIC_ERROR, {"HY000", "General error", SQL_ERROR}}, + {EStatus::TIMEOUT, {"HYT00", "Timeout expired", SQL_ERROR}}, + {EStatus::BAD_SESSION, {"08003", "Connection does not exist", SQL_ERROR}}, + {EStatus::PRECONDITION_FAILED, {"23000", "Integrity constraint violation", SQL_ERROR}}, + {EStatus::ALREADY_EXISTS, {"23000", "Integrity constraint violation", SQL_ERROR}}, + {EStatus::NOT_FOUND, {"02000", "No data found", SQL_NO_DATA}}, + {EStatus::SESSION_EXPIRED, {"08003", "Connection does not exist", SQL_ERROR}}, + {EStatus::CANCELLED, {"HY008", "Operation canceled", SQL_ERROR}}, + {EStatus::UNDETERMINED, {"40003", "Statement completion unknown", SQL_ERROR}}, + {EStatus::UNSUPPORTED, {"HYC00", "Optional feature not implemented", SQL_ERROR}}, + {EStatus::SESSION_BUSY, {"HY000", "General error - session busy", SQL_ERROR}}, + // Transport errors + {EStatus::TRANSPORT_UNAVAILABLE, {"08S01", "Communication link failure", SQL_ERROR}}, + {EStatus::CLIENT_RESOURCE_EXHAUSTED, {"HY000", "General error - resource exhausted", SQL_ERROR}}, + {EStatus::CLIENT_DEADLINE_EXCEEDED, {"HYT00", "Timeout expired", SQL_ERROR}}, + {EStatus::CLIENT_INTERNAL_ERROR, {"HY000", "General error", SQL_ERROR}}, + {EStatus::CLIENT_CANCELLED, {"HY008", "Operation canceled", SQL_ERROR}}, + {EStatus::CLIENT_UNAUTHENTICATED, {"28000", "Invalid authorization specification", SQL_ERROR}}, + {EStatus::CLIENT_LIMITS_REACHED, {"HY000", "General error - limits reached", SQL_ERROR}}, + {EStatus::CLIENT_DISCOVERY_FAILED, {"08001", "Client unable to establish connection", SQL_ERROR}}, + {EStatus::CLIENT_CALL_UNIMPLEMENTED, {"HYC00", "Optional feature not implemented", SQL_ERROR}}, + {EStatus::CLIENT_OUT_OF_RANGE, {"22003", "Numeric value out of range", SQL_ERROR}}, + }; + + const OdbcErrorMapping DEFAULT_ERROR_MAPPING = {"HY000", "Unknown YDB error", SQL_ERROR}; + + OdbcErrorMapping GetErrorMappingForStatus(EStatus status) { + auto it = ERROR_MAPPINGS.find(status); + if (it != ERROR_MAPPINGS.end()) { + return it->second; + } + return DEFAULT_ERROR_MAPPING; + } + + SQLRETURN WriteDiagCStr( + const std::string& str, + SQLPOINTER diagInfoPtr, + SQLSMALLINT bufferLength, + SQLSMALLINT* stringLengthPtr, + bool sqlStateField = false) { + std::string storage; + const std::string* src = &str; + if (sqlStateField) { + storage = str; + if (storage.size() < 5) { + storage.append(5U - storage.size(), ' '); + } else { + storage.resize(5U); + } + src = &storage; + } + const size_t fullLen = src->size(); + if (stringLengthPtr) { + *stringLengthPtr = static_cast( + std::min(fullLen, static_cast(std::numeric_limits::max()))); + } + if (!diagInfoPtr) { + return SQL_SUCCESS; + } + if (bufferLength < 0) { + return SQL_ERROR; + } + if (bufferLength == 0) { + return fullLen == 0 ? SQL_SUCCESS : SQL_SUCCESS_WITH_INFO; + } + auto* out = static_cast(diagInfoPtr); + const size_t maxData = static_cast(bufferLength - 1U); + const size_t copyLen = std::min(fullLen, maxData); + std::memcpy(out, src->data(), copyLen); + out[copyLen] = 0; + return (fullLen > maxData) ? SQL_SUCCESS_WITH_INFO : SQL_SUCCESS; + } + +} // namespace + +SQLRETURN TErrorManager::AddError(const std::string& sqlState, SQLINTEGER nativeError, const std::string& message, SQLRETURN returnCode) { + Errors_.push_back({sqlState, nativeError, message, returnCode}); + LastReturnCode_ = returnCode; + return returnCode; +} + +SQLRETURN TErrorManager::AddError(const TOdbcException& ex) { + Errors_.push_back({ex.GetSqlState(), ex.GetNativeError(), ex.GetMessage(), ex.GetReturnCode()}); + LastReturnCode_ = ex.GetReturnCode(); + return ex.GetReturnCode(); +} + +SQLRETURN TErrorManager::AddError(const TStatus& status) { + auto mapping = GetErrorMappingForStatus(status.GetStatus()); + std::string message = mapping.description; + if (!status.GetIssues().Empty()) { + message += ": " + status.GetIssues().ToString(); + } + Errors_.push_back({mapping.sqlState, static_cast(status.GetStatus()), message, mapping.returnCode}); + LastReturnCode_ = mapping.returnCode; + return mapping.returnCode; +} + +void TErrorManager::ClearErrors() { + Errors_.clear(); +} + +SQLRETURN TErrorManager::GetDiagRec(SQLSMALLINT recNumber, SQLCHAR* sqlState, SQLINTEGER* nativeError, + SQLCHAR* messageText, SQLSMALLINT bufferLength, SQLSMALLINT* textLength) { + if (recNumber < 1 || recNumber > (SQLSMALLINT)Errors_.size()) { + return SQL_NO_DATA; + } + + const auto& err = Errors_[recNumber-1]; + + if (sqlState) { + WriteDiagCStr(err.SqlState, sqlState, 6, nullptr, true); + } + + if (nativeError) { + *nativeError = err.NativeError; + } + + return WriteDiagCStr(err.Message, messageText, bufferLength, textLength, false); +} + +SQLRETURN TErrorManager::GetDiagField(SQLSMALLINT recNumber, SQLSMALLINT diagIdentifier, SQLPOINTER diagInfoPtr, + SQLSMALLINT bufferLength, SQLSMALLINT* stringLengthPtr) { + const SQLSMALLINT count = static_cast(Errors_.size()); + if (diagInfoPtr == nullptr) { + return SQL_ERROR; + } + if (recNumber == 0) { + switch (diagIdentifier) { + case SQL_DIAG_RETURNCODE: + *static_cast(diagInfoPtr) = LastReturnCode_; + return SQL_SUCCESS; + case SQL_DIAG_NUMBER: { + *static_cast(diagInfoPtr) = static_cast(count); + return SQL_SUCCESS; + } + case SQL_DIAG_ROW_COUNT: + return SQL_ERROR; + default: + return SQL_ERROR; + } + } + + if (recNumber < 1 || recNumber > count) { + return SQL_NO_DATA; + } + + const auto& err = Errors_[recNumber - 1]; + switch (diagIdentifier) { + case SQL_DIAG_SQLSTATE: + return WriteDiagCStr(err.SqlState, diagInfoPtr, bufferLength, stringLengthPtr, true); + case SQL_DIAG_NATIVE: { + *static_cast(diagInfoPtr) = err.NativeError; + return SQL_SUCCESS; + } + case SQL_DIAG_MESSAGE_TEXT: + return WriteDiagCStr(err.Message, diagInfoPtr, bufferLength, stringLengthPtr); + case SQL_DIAG_CLASS_ORIGIN: + return WriteDiagCStr("ODBC 3.0", diagInfoPtr, bufferLength, stringLengthPtr); + case SQL_DIAG_SUBCLASS_ORIGIN: + return WriteDiagCStr("ODBC 3.0", diagInfoPtr, bufferLength, stringLengthPtr); + case SQL_DIAG_CONNECTION_NAME: + case SQL_DIAG_SERVER_NAME: + return WriteDiagCStr("", diagInfoPtr, bufferLength, stringLengthPtr); + case SQL_DIAG_COLUMN_NUMBER: + *static_cast(diagInfoPtr) = SQL_COLUMN_NUMBER_UNKNOWN; + return SQL_SUCCESS; + case SQL_DIAG_ROW_NUMBER: + *static_cast(diagInfoPtr) = SQL_ROW_NUMBER_UNKNOWN; + return SQL_SUCCESS; + default: + return SQL_ERROR; + } +} + +SQLRETURN HandleOdbcExceptions( + SQLHANDLE handlePtr, + std::function&& func, + ENullInputHandlePolicy nullInputPolicy) { + if (!handlePtr && nullInputPolicy != ENullInputHandlePolicy::Allow) { + return SQL_INVALID_HANDLE; + } + + try { + const SQLRETURN r = func(); + if (handlePtr) { + static_cast(handlePtr)->SetLastReturnCode(r); + } + return r; + } catch (...) { + if (handlePtr) { + static_cast(handlePtr)->SetLastReturnCode(SQL_ERROR); + } + return SQL_ERROR; + } +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/error_manager.h b/odbc/src/utils/error_manager.h new file mode 100644 index 00000000000..6677f3fdcdf --- /dev/null +++ b/odbc/src/utils/error_manager.h @@ -0,0 +1,157 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +#include + +namespace NYdb::NOdbc { + +struct TErrorInfo { + std::string SqlState; + SQLINTEGER NativeError; + std::string Message; + SQLRETURN ReturnCode; +}; + +using TErrorList = std::vector; + +class TOdbcException : public std::exception { +public: + TOdbcException(const std::string& sqlState, SQLINTEGER nativeError, + const std::string& message, SQLRETURN returnCode = SQL_ERROR) + : SqlState_(sqlState) + , NativeError_(nativeError) + , Message_(message) + , ReturnCode_(returnCode) + {} + + const std::string& GetSqlState() const { + return SqlState_; + } + + SQLINTEGER GetNativeError() const { + return NativeError_; + } + + const std::string& GetMessage() const { + return Message_; + } + + SQLRETURN GetReturnCode() const { + return ReturnCode_; + } + + const char* what() const noexcept override { + return Message_.c_str(); + } + +private: + std::string SqlState_; + SQLINTEGER NativeError_; + std::string Message_; + SQLRETURN ReturnCode_; +}; + +class TErrorManager { +public: + SQLRETURN AddError(const std::string& sqlState, SQLINTEGER nativeError, const std::string& message, SQLRETURN returnCode = SQL_ERROR); + SQLRETURN AddError(const TOdbcException& ex); + SQLRETURN AddError(const TStatus& status); + + void ClearErrors(); + std::recursive_mutex& GetMutex() const noexcept { return Mutex_; } + + void SetLastReturnCode(SQLRETURN code) { + LastReturnCode_ = code; + } + [[nodiscard]] SQLRETURN GetLastReturnCode() const { + return LastReturnCode_; + } + + SQLRETURN GetDiagRec(SQLSMALLINT recNumber, SQLCHAR* sqlState, SQLINTEGER* nativeError, + SQLCHAR* messageText, SQLSMALLINT bufferLength, SQLSMALLINT* textLength); + virtual SQLRETURN GetDiagField(SQLSMALLINT recNumber, SQLSMALLINT diagIdentifier, + SQLPOINTER diagInfoPtr, SQLSMALLINT bufferLength, SQLSMALLINT* stringLengthPtr); + +private: + mutable std::recursive_mutex Mutex_; + TErrorList Errors_; + SQLRETURN LastReturnCode_ = SQL_SUCCESS; +}; + +enum class ENullInputHandlePolicy : unsigned char { + Reject, + Allow, +}; + +template +SQLRETURN HandleOdbcExceptionsConsuming(SQLHANDLE handlePtr, std::function&& func) { + if (!handlePtr) { + return SQL_INVALID_HANDLE; + } + auto handle = static_cast(handlePtr); + handle->ClearErrors(); + + try { + return func(handle); + } catch (const NStatusHelpers::TYdbErrorException& ex) { + return handle->AddError(ex.GetStatus()); + } catch (const TOdbcException& ex) { + return handle->AddError(ex); + } catch (const std::exception& ex) { + return handle->AddError("HY000", 0, ex.what()); + } catch (...) { + return handle->AddError("HY000", 0, "Unknown error"); + } +} + +template +SQLRETURN HandleOdbcDiagnostics(SQLHANDLE handlePtr, std::function&& func) { + if (!handlePtr) { + return SQL_INVALID_HANDLE; + } + auto* handle = static_cast(handlePtr); + std::lock_guard lock(handle->GetMutex()); + try { + return func(handle); + } catch (...) { + return SQL_ERROR; + } +} + +template +SQLRETURN HandleOdbcExceptions(SQLHANDLE handlePtr, std::function&& func) { + if (!handlePtr) { + return SQL_INVALID_HANDLE; + } + auto handle = static_cast(handlePtr); + std::lock_guard lock(handle->GetMutex()); + handle->ClearErrors(); + + try { + const SQLRETURN ret = func(handle); + handle->SetLastReturnCode(ret); + return ret; + } catch (const NStatusHelpers::TYdbErrorException& ex) { + return handle->AddError(ex.GetStatus()); + } catch (const TOdbcException& ex) { + return handle->AddError(ex); + } catch (const std::exception& ex) { + return handle->AddError("HY000", 0, ex.what()); + } catch (...) { + return handle->AddError("HY000", 0, "Unknown error"); + } +} + +SQLRETURN HandleOdbcExceptions( + SQLHANDLE handlePtr, + std::function&& func, + ENullInputHandlePolicy nullInputPolicy = ENullInputHandlePolicy::Reject); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/escape.cpp b/odbc/src/utils/escape.cpp new file mode 100644 index 00000000000..f6ab7dcf9b6 --- /dev/null +++ b/odbc/src/utils/escape.cpp @@ -0,0 +1,395 @@ +#include "escape.h" +#include "sql_type_map.h" + +#include +#include +#include +#include + +namespace NYdb::NOdbc { +namespace { + +bool EqualNoCase(std::string_view lhs, std::string_view rhs) { + return lhs.size() == rhs.size() && + std::equal(lhs.begin(), lhs.end(), rhs.begin(), [](char leftCh, char rightCh) { + return std::tolower(static_cast(leftCh)) == + std::tolower(static_cast(rightCh)); + }); +} + +void SkipLeadingWhitespace(std::string_view sql, size_t& cursor) { + const auto strEnd = sql.end(); + const auto firstNonSpace = std::find_if_not( + sql.begin() + static_cast(cursor), + strEnd, + [](unsigned char byte) { + return std::isspace(byte) != 0; + }); + cursor = static_cast(firstNonSpace - sql.begin()); +} + +bool ReadIdent(std::string_view sql, size_t& cursor, std::string_view* outIdent) { + SkipLeadingWhitespace(sql, cursor); + const size_t identStart = cursor; + const auto afterIdent = std::find_if_not( + sql.begin() + static_cast(cursor), + sql.end(), + [](unsigned char byte) { + return std::isalpha(byte) != 0 || byte == '_'; + }); + cursor = static_cast(afterIdent - sql.begin()); + if (cursor == identStart) { + return false; + } + *outIdent = std::string_view(sql.data() + identStart, cursor - identStart); + return true; +} + +bool ParseSingleQuoted(std::string_view sql, size_t& cursor, std::string* outValue) { + SkipLeadingWhitespace(sql, cursor); + if (cursor >= sql.size() || sql[cursor] != '\'') { + return false; + } + ++cursor; + outValue->clear(); + while (cursor < sql.size()) { + if (sql[cursor] == '\'') { + if (cursor + 1 < sql.size() && sql[cursor + 1] == '\'') { + outValue->push_back('\''); + cursor += 2; + continue; + } + ++cursor; + return true; + } + outValue->push_back(sql[cursor++]); + } + return false; +} + +size_t FindMatchingCloseBrace(std::string_view sql, size_t openBrace) { + if (openBrace >= sql.size() || sql[openBrace] != '{') { + return std::string_view::npos; + } + int braceDepth = 1; + for (size_t idx = openBrace + 1; idx < sql.size(); ++idx) { + if (sql[idx] == '{') { + ++braceDepth; + } else if (sql[idx] == '}') { + --braceDepth; + if (braceDepth == 0) { + return idx; + } + } + } + return std::string_view::npos; +} + +std::string NormalizeOdbcTimestampLiteral(const std::string& raw) { + std::string normalized = raw; + const auto firstSpace = std::find(normalized.begin(), normalized.end(), ' '); + if (firstSpace != normalized.end()) { + *firstSpace = 'T'; + } + if (std::find(normalized.begin(), normalized.end(), 'Z') == normalized.end()) { + normalized.push_back('Z'); + } + return normalized; +} + +std::string RewriteOdbcEscapesImpl(std::string_view sql); + + +enum class OdbcBraceKind { + OutputProcedureCall, // {?= call ... } + FnBody, // {fn ...} + OjBody, // {oj ...} + DateLiteral, // {d '...'} + TimeLiteral, // {t '...'} + TimestampLiteral, // {ts '...'} + ProcedureCall, // {call ...} + LikeEscape, // {escape '...'} +}; + +struct OdbcBraceParsed { + OdbcBraceKind Kind; + std::string_view RecurseTail; + std::string QuotedValue; +}; + +std::optional TryParseOutputCallBrace(std::string_view sql, size_t parsePos, size_t closeBrace) { + if (parsePos + 1 >= sql.size() || sql[parsePos] != '?' || sql[parsePos + 1] != '=') { + return std::nullopt; + } + size_t inner = parsePos + 2; + SkipLeadingWhitespace(sql, inner); + std::string_view keyword; + if (!ReadIdent(sql, inner, &keyword) || !EqualNoCase(keyword, "call")) { + return std::nullopt; + } + SkipLeadingWhitespace(sql, inner); + if (inner > closeBrace) { + return std::nullopt; + } + OdbcBraceParsed parsed; + parsed.Kind = OdbcBraceKind::OutputProcedureCall; + parsed.RecurseTail = std::string_view(sql.data() + inner, closeBrace - inner); + return parsed; +} + +std::optional MakeRecurseTailBrace(OdbcBraceKind kind, std::string_view sql, size_t& parsePos, size_t closeBrace) { + SkipLeadingWhitespace(sql, parsePos); + if (parsePos > closeBrace) { + return std::nullopt; + } + OdbcBraceParsed parsed; + parsed.Kind = kind; + parsed.RecurseTail = std::string_view(sql.data() + parsePos, closeBrace - parsePos); + return parsed; +} + +std::optional MakeQuotedBrace(OdbcBraceKind kind, std::string_view sql, size_t& parsePos, size_t closeBrace) { + std::string quotedLit; + if (!ParseSingleQuoted(sql, parsePos, "edLit) || parsePos > closeBrace) { + return std::nullopt; + } + SkipLeadingWhitespace(sql, parsePos); + if (parsePos != closeBrace) { + return std::nullopt; + } + OdbcBraceParsed parsed; + parsed.Kind = kind; + parsed.QuotedValue = std::move(quotedLit); + return parsed; +} + +struct BraceKeywordSpec { + std::string_view Keyword; + OdbcBraceKind Kind; + bool IsQuotedLiteral; +}; + +static constexpr BraceKeywordSpec kBraceKeywordSpecs[] = { + {"fn", OdbcBraceKind::FnBody, false}, + {"oj", OdbcBraceKind::OjBody, false}, + {"d", OdbcBraceKind::DateLiteral, true}, + {"t", OdbcBraceKind::TimeLiteral, true}, + {"ts", OdbcBraceKind::TimestampLiteral, true}, + {"call", OdbcBraceKind::ProcedureCall, false}, + {"escape", OdbcBraceKind::LikeEscape, true}, +}; + +std::optional TryParseOdbcBrace(std::string_view sql, size_t openBrace, size_t closeBrace) { + size_t parsePos = openBrace + 1; + SkipLeadingWhitespace(sql, parsePos); + + if (std::optional outputCall = TryParseOutputCallBrace(sql, parsePos, closeBrace)) { + return outputCall; + } + if (parsePos + 1 < sql.size() && sql[parsePos] == '?' && sql[parsePos + 1] == '=') { + return std::nullopt; + } + + std::string_view token; + if (!ReadIdent(sql, parsePos, &token)) { + return std::nullopt; + } + + for (const BraceKeywordSpec& spec : kBraceKeywordSpecs) { + if (!EqualNoCase(token, spec.Keyword)) { + continue; + } + if (spec.IsQuotedLiteral) { + return MakeQuotedBrace(spec.Kind, sql, parsePos, closeBrace); + } + return MakeRecurseTailBrace(spec.Kind, sql, parsePos, closeBrace); + } + + return std::nullopt; +} + +void AppendRewrittenBrace(std::string& rewritten, const OdbcBraceParsed& parsed) { + switch (parsed.Kind) { + case OdbcBraceKind::OutputProcedureCall: + case OdbcBraceKind::ProcedureCall: + rewritten += "CALL "; + rewritten.append(RewriteOdbcEscapesImpl(parsed.RecurseTail)); + return; + case OdbcBraceKind::FnBody: + case OdbcBraceKind::OjBody: + rewritten.append(RewriteOdbcEscapesImpl(parsed.RecurseTail)); + return; + case OdbcBraceKind::DateLiteral: + rewritten += "CAST('"; + rewritten += parsed.QuotedValue; + rewritten += "' AS Date)"; + return; + case OdbcBraceKind::TimeLiteral: + rewritten += "CAST('"; + rewritten += parsed.QuotedValue; + rewritten += "' AS Time)"; + return; + case OdbcBraceKind::TimestampLiteral: { + const std::string normalizedTs = NormalizeOdbcTimestampLiteral(parsed.QuotedValue); + rewritten += "CAST('"; + rewritten += normalizedTs; + rewritten += "' AS Datetime)"; + return; + } + case OdbcBraceKind::LikeEscape: + rewritten += " ESCAPE '"; + rewritten += parsed.QuotedValue; + rewritten += '\''; + return; + } +} + +std::string RewriteOdbcEscapesImpl(std::string_view sql) { + std::string rewritten; + rewritten.reserve(sql.size()); + + for (size_t readPos = 0; readPos < sql.size();) { + if (sql[readPos] != '{') { + rewritten.push_back(sql[readPos++]); + continue; + } + + const size_t closeBrace = FindMatchingCloseBrace(sql, readPos); + if (closeBrace == std::string_view::npos) { + rewritten.push_back(sql[readPos++]); + continue; + } + + if (std::optional parsedBrace = TryParseOdbcBrace(sql, readPos, closeBrace)) { + AppendRewrittenBrace(rewritten, *parsedBrace); + readPos = closeBrace + 1; + continue; + } + + rewritten.push_back(sql[readPos++]); + } + + return rewritten; +} + +std::string RewriteOdbcConvertCalls(std::string_view sql); + +class TOdbcConvertCallRewriter { +public: + explicit TOdbcConvertCallRewriter(std::string_view sql) + : Sql_(sql) { + Rewritten_.reserve(sql.size()); + } + + std::string TakeResult() && { + return std::move(Rewritten_); + } + + void Run() { + while (SegmentStart_ < Sql_.size()) { + const std::optional convertKeywordPos = FindNextConvertKeyword(SegmentStart_); + if (!convertKeywordPos) { + Rewritten_.append(Sql_.substr(SegmentStart_)); + break; + } + Rewritten_.append(Sql_.substr(SegmentStart_, *convertKeywordPos - SegmentStart_)); + if (!TryRewriteConvertAt(*convertKeywordPos)) { + break; + } + } + } + +private: + static constexpr size_t kConvertTokenLen = 7; + + std::optional FindNextConvertKeyword(size_t from) const { + for (size_t probePos = from; probePos + kConvertTokenLen <= Sql_.size(); ++probePos) { + if (!EqualNoCase(Sql_.substr(probePos, kConvertTokenLen), "CONVERT")) { + continue; + } + size_t afterKeyword = probePos + kConvertTokenLen; + SkipLeadingWhitespace(Sql_, afterKeyword); + if (afterKeyword < Sql_.size() && Sql_[afterKeyword] == '(') { + return probePos; + } + } + return std::nullopt; + } + + bool TryRewriteConvertAt(size_t convertKeywordPos) { + size_t parsePos = convertKeywordPos + kConvertTokenLen; + SkipLeadingWhitespace(Sql_, parsePos); + if (parsePos >= Sql_.size() || Sql_[parsePos] != '(') { + Rewritten_.append(Sql_.substr(convertKeywordPos, kConvertTokenLen)); + SegmentStart_ = convertKeywordPos + kConvertTokenLen; + return true; + } + ++parsePos; + + int parenDepth = 1; + const size_t firstArgStart = parsePos; + std::optional typeCommaPos; + for (; parsePos < Sql_.size(); ++parsePos) { + if (Sql_[parsePos] == '(') { + ++parenDepth; + } else if (Sql_[parsePos] == ')') { + --parenDepth; + } else if (Sql_[parsePos] == ',' && parenDepth == 1) { + typeCommaPos = parsePos; + break; + } + } + if (!typeCommaPos) { + Rewritten_.append(Sql_.substr(convertKeywordPos)); + return false; + } + + const std::string_view firstArg(Sql_.data() + firstArgStart, *typeCommaPos - firstArgStart); + parsePos = *typeCommaPos + 1; + SkipLeadingWhitespace(Sql_, parsePos); + const size_t sqlTypeStart = parsePos; + const auto sqlTypeEnd = std::find_if_not( + Sql_.begin() + static_cast(parsePos), + Sql_.end(), + [](unsigned char byte) { + return std::isalpha(byte) != 0 || byte == '_'; + }); + parsePos = static_cast(sqlTypeEnd - Sql_.begin()); + const std::string_view sqlTypeToken(Sql_.data() + sqlTypeStart, parsePos - sqlTypeStart); + SkipLeadingWhitespace(Sql_, parsePos); + if (parsePos >= Sql_.size() || Sql_[parsePos] != ')') { + Rewritten_.append(Sql_.substr(convertKeywordPos)); + return false; + } + + const std::string yqlType = MapSqlTypeToken(sqlTypeToken); + Rewritten_ += "CAST("; + Rewritten_ += RewriteOdbcConvertCalls(RewriteOdbcEscapesImpl(firstArg)); + Rewritten_ += " AS "; + Rewritten_ += yqlType; + Rewritten_ += ')'; + SegmentStart_ = parsePos + 1; + return true; + } + + std::string_view Sql_; + std::string Rewritten_; + size_t SegmentStart_ = 0; +}; + +std::string RewriteOdbcConvertCalls(std::string_view sql) { + TOdbcConvertCallRewriter rewriter(sql); + rewriter.Run(); + return std::move(rewriter).TakeResult(); +} + +} // namespace + + + +std::string RewriteOdbcEscapes(const std::string& sql) { + std::string afterBraceRewrite = RewriteOdbcEscapesImpl(sql); + return RewriteOdbcConvertCalls(afterBraceRewrite); +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/escape.h b/odbc/src/utils/escape.h new file mode 100644 index 00000000000..7397a128450 --- /dev/null +++ b/odbc/src/utils/escape.h @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace NYdb::NOdbc { + +std::string RewriteOdbcEscapes(const std::string& sql); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/param_rewrite.cpp b/odbc/src/utils/param_rewrite.cpp new file mode 100644 index 00000000000..315e4c96195 --- /dev/null +++ b/odbc/src/utils/param_rewrite.cpp @@ -0,0 +1,165 @@ +#include "param_rewrite.h" +#include "sql_type_map.h" + +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +namespace { + +bool IsParamMark(std::string_view sql, size_t i) { + if (sql[i] != '?') { + return false; + } + if (i > 0) { + const char prev = sql[i - 1]; + if (std::isalnum(static_cast(prev)) || prev == '_' || prev == ')') { + return false; + } + } + return true; +} + +bool TryParseDollarParam(std::string_view sql, size_t i, SQLUSMALLINT& index) { + if (sql.size() < i + 3 || sql[i] != '$' || sql[i + 1] != 'p' || !std::isdigit(static_cast(sql[i + 2]))) { + return false; + } + unsigned n = 0; + for (size_t j = i + 2; j < sql.size() && std::isdigit(static_cast(sql[j])); ++j) { + n = n * 10 + static_cast(sql[j] - '0'); + } + index = static_cast(n); + return true; +} + +} // namespace + +SQLSMALLINT CountOdbcParams(std::string_view sql) { + SQLSMALLINT questionMarkCount = 0; + SQLSMALLINT maxDollarIndex = 0; + bool inQuote = false; + size_t braceDepth = 0; + + for (size_t i = 0; i < sql.size(); ++i) { + const char ch = sql[i]; + if (inQuote) { + if (ch == '\'' && i + 1 < sql.size() && sql[i + 1] == '\'') { + ++i; + } else if (ch == '\'') { + inQuote = false; + } + continue; + } + if (ch == '\'') { + inQuote = true; + continue; + } + if (ch == '{') { + ++braceDepth; + continue; + } + if (ch == '}' && braceDepth > 0) { + --braceDepth; + continue; + } + if (braceDepth == 0) { + if (IsParamMark(sql, i)) { + ++questionMarkCount; + continue; + } + SQLUSMALLINT index = 0; + if (TryParseDollarParam(sql, i, index)) { + maxDollarIndex = std::max(maxDollarIndex, static_cast(index)); + } + } + } + + if (questionMarkCount > 0) { + return questionMarkCount; + } + return maxDollarIndex; +} + +TParamRewriteResult RewriteOdbcQuestionMarks( + std::string_view sql, + const std::vector& boundParams) { + std::string body; + body.reserve(sql.size()); + size_t questionMarkCount = 0; + bool inQuote = false; + size_t braceDepth = 0; + std::set paramIndices; + + for (size_t i = 0; i < sql.size(); ++i) { + const char ch = sql[i]; + if (inQuote) { + body.push_back(ch); + if (ch == '\'' && i + 1 < sql.size() && sql[i + 1] == '\'') { + body.push_back('\''); + ++i; + } else if (ch == '\'') { + inQuote = false; + } + continue; + } + if (ch == '\'') { + inQuote = true; + body.push_back(ch); + continue; + } + if (ch == '{') { + ++braceDepth; + body.push_back(ch); + continue; + } + if (ch == '}' && braceDepth > 0) { + --braceDepth; + body.push_back(ch); + continue; + } + if (braceDepth == 0) { + if (IsParamMark(sql, i)) { + const auto index = static_cast(++questionMarkCount); + paramIndices.insert(index); + body += "$p"; + body += std::to_string(index); + continue; + } + SQLUSMALLINT index = 0; + if (TryParseDollarParam(sql, i, index)) { + paramIndices.insert(index); + } + } + body.push_back(ch); + } + + if (paramIndices.empty()) { + return {.Sql = std::string(sql)}; + } + if (questionMarkCount > 0 && questionMarkCount != boundParams.size()) { + return {.Success = false, .SqlState = "07002", .Message = "COUNT field incorrect"}; + } + + std::string declares; + for (const SQLUSMALLINT index : paramIndices) { + const std::string declare = "DECLARE $p" + std::to_string(index) + " AS"; + if (sql.find(declare) != std::string_view::npos) { + continue; + } + const auto bound = std::ranges::find(boundParams, index, &TBoundParam::ParamNumber); + if (bound == boundParams.end()) { + return {.Success = false, .SqlState = "07002", .Message = "COUNT field incorrect"}; + } + declares += declare + " " + FormatYqlParamDeclareType(bound->ParameterType) + ";\n"; + } + if (declares.empty()) { + return {.Sql = body}; + } + return {.Sql = declares + body}; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/param_rewrite.h b/odbc/src/utils/param_rewrite.h new file mode 100644 index 00000000000..f1b41d51218 --- /dev/null +++ b/odbc/src/utils/param_rewrite.h @@ -0,0 +1,24 @@ +#pragma once + +#include "bindings.h" + +#include +#include +#include + +namespace NYdb::NOdbc { + +struct TParamRewriteResult { + std::string Sql; + bool Success = true; + std::string SqlState; + std::string Message; +}; + +TParamRewriteResult RewriteOdbcQuestionMarks( + std::string_view sql, + const std::vector& boundParams); + +SQLSMALLINT CountOdbcParams(std::string_view sql); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/sql_like.h b/odbc/src/utils/sql_like.h new file mode 100644 index 00000000000..f51c10ca28c --- /dev/null +++ b/odbc/src/utils/sql_like.h @@ -0,0 +1,49 @@ +#pragma once + +#include + +namespace NYdb::NOdbc { + +// SQL LIKE — '%' is any substring, '_' is any single character. +inline bool SqlLikeMatch(std::string_view text, std::string_view pattern) { + size_t textPos = 0; + size_t patPos = 0; + size_t lastPercentPat = std::string_view::npos; + size_t textStartAfterPercent = 0; + + const size_t textLen = text.size(); + const size_t patLen = pattern.size(); + + while (textPos < textLen) { + const bool morePat = patPos < patLen; + const char patCh = morePat ? pattern[patPos] : '\0'; + + if (morePat && patCh != '%' && (patCh == '_' || patCh == text[textPos])) { + ++textPos; + ++patPos; + continue; + } + + if (morePat && patCh == '%') { + lastPercentPat = patPos++; + textStartAfterPercent = textPos; + continue; + } + + if (lastPercentPat != std::string_view::npos) { + patPos = lastPercentPat + 1; + ++textStartAfterPercent; + textPos = textStartAfterPercent; + continue; + } + + return false; + } + + while (patPos < patLen && pattern[patPos] == '%') { + ++patPos; + } + return patPos == patLen; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/sql_type_map.cpp b/odbc/src/utils/sql_type_map.cpp new file mode 100644 index 00000000000..c8ba5f19af3 --- /dev/null +++ b/odbc/src/utils/sql_type_map.cpp @@ -0,0 +1,68 @@ +#include "sql_type_map.h" + +#include +#include +#include +#include + +namespace NYdb::NOdbc { +namespace { + +constexpr std::array TypeSpecs{ + TSqlTypeSpec{SQL_BIGINT, "BIGINT", "Int64", 19, true}, + TSqlTypeSpec{SQL_INTEGER, "INTEGER", "Int32", 10, true}, + TSqlTypeSpec{SQL_SMALLINT, "SMALLINT", "Int16", 5, true}, + TSqlTypeSpec{SQL_DOUBLE, "DOUBLE", "Double", 15, true}, + TSqlTypeSpec{SQL_REAL, "REAL", "Float", 7, true}, + TSqlTypeSpec{SQL_VARCHAR, "VARCHAR", "Utf8", 255, true}, + TSqlTypeSpec{SQL_CHAR, "CHAR", "Utf8", 255, true}, + TSqlTypeSpec{SQL_LONGVARCHAR, "LONGVARCHAR", "Utf8", 4096, false}, + TSqlTypeSpec{SQL_WCHAR, "WCHAR", "Utf8", 255, false}, + TSqlTypeSpec{SQL_WVARCHAR, "WVARCHAR", "Utf8", 255, false}, + TSqlTypeSpec{SQL_WLONGVARCHAR, "WLONGVARCHAR", "Utf8", 4096, false}, + TSqlTypeSpec{SQL_BIT, "BIT", "Bool", 1, false}, + TSqlTypeSpec{SQL_TINYINT, "TINYINT", "Int8", 3, false}, + TSqlTypeSpec{SQL_FLOAT, "FLOAT", "Double", 15, false}, + TSqlTypeSpec{SQL_DECIMAL, "DECIMAL", "Decimal(22, 9)", 22, false}, + TSqlTypeSpec{SQL_NUMERIC, "NUMERIC", "Decimal(22, 9)", 22, false}, + TSqlTypeSpec{SQL_BINARY, "BINARY", "String", 4096, false}, + TSqlTypeSpec{SQL_VARBINARY, "VARBINARY", "String", 4096, false}, + TSqlTypeSpec{SQL_LONGVARBINARY, "LONGVARBINARY", "String", 4096, false}, + TSqlTypeSpec{SQL_TYPE_DATE, "DATE", "Date", 10, false}, + TSqlTypeSpec{SQL_TYPE_TIME, "TIME", "Time", 8, false}, + TSqlTypeSpec{SQL_TYPE_TIMESTAMP, "TIMESTAMP", "Datetime", 26, false}, +}; + +std::string ToUpperAscii(std::string_view value) { + std::string upper(value); + std::ranges::transform(upper, upper.begin(), [](unsigned char byte) { + return static_cast(std::toupper(byte)); + }); + return upper; +} + +} // namespace + +std::span GetSqlTypeSpecs() { + return TypeSpecs; +} + +const TSqlTypeSpec* FindSqlTypeSpec(SQLSMALLINT sqlType) { + const auto it = std::ranges::find(TypeSpecs, sqlType, &TSqlTypeSpec::Type); + return it == TypeSpecs.end() ? nullptr : &*it; +} + +std::string MapSqlTypeToken(std::string_view sqlType) { + std::string key = ToUpperAscii(sqlType); + if (key.starts_with("SQL_")) key.erase(0, 4); + if (key.starts_with("TYPE_")) key.erase(0, 5); + const auto it = std::ranges::find(TypeSpecs, key, &TSqlTypeSpec::Name); + return it == TypeSpecs.end() ? key : std::string(it->YqlType); +} + +std::string FormatYqlParamDeclareType(SQLSMALLINT sqlType) { + const TSqlTypeSpec* spec = FindSqlTypeSpec(sqlType); + return (spec ? std::string(spec->YqlType) : std::to_string(sqlType)) + '?'; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/sql_type_map.h b/odbc/src/utils/sql_type_map.h new file mode 100644 index 00000000000..913064babef --- /dev/null +++ b/odbc/src/utils/sql_type_map.h @@ -0,0 +1,25 @@ +#pragma once + +#include +#include + +#include +#include +#include + +namespace NYdb::NOdbc { + +struct TSqlTypeSpec { + SQLSMALLINT Type; + std::string_view Name; + std::string_view YqlType; + SQLULEN ColumnSize; + bool Advertise; +}; + +std::span GetSqlTypeSpecs(); +const TSqlTypeSpec* FindSqlTypeSpec(SQLSMALLINT sqlType); +std::string MapSqlTypeToken(std::string_view sqlType); +std::string FormatYqlParamDeclareType(SQLSMALLINT sqlType); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/status_util.cpp b/odbc/src/utils/status_util.cpp new file mode 100644 index 00000000000..99620babd88 --- /dev/null +++ b/odbc/src/utils/status_util.cpp @@ -0,0 +1,11 @@ +#include "status_util.h" + +#include + +namespace NYdb::NOdbc { + +NYdb::TStatus StatusFrom(const NYdb::TStatus& ydbStatus) { + return NYdb::TStatus(ydbStatus.GetStatus(), NYdb::NIssue::TIssues(ydbStatus.GetIssues())); +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/status_util.h b/odbc/src/utils/status_util.h new file mode 100644 index 00000000000..43595c69ac1 --- /dev/null +++ b/odbc/src/utils/status_util.h @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace NYdb::NOdbc { + +NYdb::TStatus StatusFrom(const NYdb::TStatus& ydbStatus); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/type_info_rows.cpp b/odbc/src/utils/type_info_rows.cpp new file mode 100644 index 00000000000..24fd54d7d2d --- /dev/null +++ b/odbc/src/utils/type_info_rows.cpp @@ -0,0 +1,66 @@ +#include "type_info_rows.h" +#include "sql_type_map.h" + +#include + +#include +#include +#include + +namespace NYdb::NOdbc { +namespace { + +TValue MakeOptionalInt16(SQLSMALLINT value) { + return TValueBuilder().OptionalInt16(value).Build(); +} + +TValue MakeOptionalInt32(SQLINTEGER value) { + return TValueBuilder().OptionalInt32(value).Build(); +} + +TValue MakeNullUtf8() { + return TValueBuilder().OptionalUtf8(std::nullopt).Build(); +} + +std::vector MakeTypeInfoRow(const TSqlTypeSpec& spec) { + std::string typeName(spec.Name); + std::ranges::transform(typeName, typeName.begin(), [](unsigned char c) { + return static_cast(std::tolower(c)); + }); + return { + TValueBuilder().Utf8(typeName).Build(), + TValueBuilder().Int16(spec.Type).Build(), + MakeOptionalInt32(static_cast(spec.ColumnSize)), + MakeNullUtf8(), + MakeNullUtf8(), + MakeNullUtf8(), + MakeOptionalInt16(SQL_NULLABLE), + MakeOptionalInt16(SQL_FALSE), + MakeOptionalInt16(SQL_PRED_SEARCHABLE), + MakeNullUtf8(), + MakeOptionalInt16(SQL_FALSE), + MakeOptionalInt16(SQL_FALSE), + TValueBuilder().OptionalUtf8(typeName).Build(), + MakeOptionalInt16(0), + MakeOptionalInt16(0), + MakeOptionalInt16(spec.Type), + MakeOptionalInt16(0), + MakeOptionalInt32(10), + MakeOptionalInt32(0), + }; +} + +} // namespace + +TTable BuildTypeInfoRows(SQLSMALLINT dataType) { + TTable table; + for (const TSqlTypeSpec& spec : GetSqlTypeSpecs()) { + if (!spec.Advertise || (dataType != SQL_ALL_TYPES && spec.Type != dataType)) { + continue; + } + table.push_back(MakeTypeInfoRow(spec)); + } + return table; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/type_info_rows.h b/odbc/src/utils/type_info_rows.h new file mode 100644 index 00000000000..23c87f6cf0e --- /dev/null +++ b/odbc/src/utils/type_info_rows.h @@ -0,0 +1,11 @@ +#pragma once + +#include "cursor.h" + +#include + +namespace NYdb::NOdbc { + +TTable BuildTypeInfoRows(SQLSMALLINT dataType); + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/types.cpp b/odbc/src/utils/types.cpp new file mode 100644 index 00000000000..d991ca24f6a --- /dev/null +++ b/odbc/src/utils/types.cpp @@ -0,0 +1,136 @@ +#include "types.h" +#include "sql_type_map.h" + +namespace NYdb { +namespace NOdbc { + +SQLSMALLINT GetTypeId(const TType& type) { + TTypeParser typeParser(type); + size_t openedOptionals = 0; + while (typeParser.GetKind() == TTypeParser::ETypeKind::Optional) { + typeParser.OpenOptional(); + ++openedOptionals; + } + + auto closeOpenedOptionals = [&]() { + while (openedOptionals > 0) { + typeParser.CloseOptional(); + --openedOptionals; + } + }; + + const auto kind = typeParser.GetKind(); + if (kind == TTypeParser::ETypeKind::Primitive) { + const auto primitive = typeParser.GetPrimitive(); + closeOpenedOptionals(); + switch (primitive) { + case EPrimitiveType::Bool: + return SQL_BIT; + case EPrimitiveType::Int8: + case EPrimitiveType::Uint8: + return SQL_TINYINT; + case EPrimitiveType::Int16: + case EPrimitiveType::Uint16: + return SQL_SMALLINT; + case EPrimitiveType::Int32: + case EPrimitiveType::Uint32: + return SQL_INTEGER; + case EPrimitiveType::Int64: + case EPrimitiveType::Uint64: + return SQL_BIGINT; + case EPrimitiveType::Float: + return SQL_REAL; + case EPrimitiveType::Double: + return SQL_DOUBLE; + case EPrimitiveType::Date: + case EPrimitiveType::Date32: + case EPrimitiveType::TzDate: + return SQL_TYPE_DATE; + case EPrimitiveType::Datetime: + case EPrimitiveType::Timestamp: + case EPrimitiveType::Datetime64: + case EPrimitiveType::Timestamp64: + case EPrimitiveType::TzDatetime: + case EPrimitiveType::TzTimestamp: + return SQL_TYPE_TIMESTAMP; + case EPrimitiveType::Interval: + case EPrimitiveType::Interval64: + return SQL_BIGINT; + case EPrimitiveType::String: + return SQL_VARBINARY; + case EPrimitiveType::Utf8: + case EPrimitiveType::Yson: + case EPrimitiveType::Json: + case EPrimitiveType::JsonDocument: + case EPrimitiveType::DyNumber: + return SQL_VARCHAR; + case EPrimitiveType::Uuid: + return SQL_GUID; + } + } + + closeOpenedOptionals(); + if (kind == TTypeParser::ETypeKind::Decimal) { + return SQL_DECIMAL; + } + return SQL_UNKNOWN_TYPE; +} + +SQLSMALLINT IsNullable(const TType& type) { + TTypeParser typeParser(type); + if (typeParser.GetKind() == TTypeParser::ETypeKind::Optional || typeParser.GetKind() == TTypeParser::ETypeKind::Null) { + return SQL_NULLABLE; + } + + return SQL_NO_NULLS; +} + +SQLULEN GetColumnSize(SQLSMALLINT sqlType) { + const TSqlTypeSpec* spec = FindSqlTypeSpec(sqlType); + return spec ? spec->ColumnSize : sqlType == SQL_GUID ? 36 : 4096; +} + +std::optional GetDecimalDigits(const TType& type) { + TTypeParser typeParser(type); + if (typeParser.GetKind() != TTypeParser::ETypeKind::Primitive) { + return std::nullopt; + } + + switch (typeParser.GetPrimitive()) { + case EPrimitiveType::Int64: + case EPrimitiveType::Uint64: + case EPrimitiveType::Int32: + case EPrimitiveType::Uint32: + case EPrimitiveType::Int16: + case EPrimitiveType::Uint16: + case EPrimitiveType::Int8: + case EPrimitiveType::Uint8: + return 0; + default: + return std::nullopt; + } +} + +std::optional GetRadix(const TType& type) { + TTypeParser typeParser(type); + if (typeParser.GetKind() != TTypeParser::ETypeKind::Primitive) { + return std::nullopt; + } + + switch (typeParser.GetPrimitive()) { + case EPrimitiveType::Int64: + case EPrimitiveType::Uint64: + case EPrimitiveType::Int32: + case EPrimitiveType::Uint32: + case EPrimitiveType::Int16: + case EPrimitiveType::Uint16: + case EPrimitiveType::Int8: + case EPrimitiveType::Uint8: + return 10; + default: + return std::nullopt; + } +} + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/types.h b/odbc/src/utils/types.h new file mode 100644 index 00000000000..9428cafebb0 --- /dev/null +++ b/odbc/src/utils/types.h @@ -0,0 +1,20 @@ +#pragma once + +#include + +#include +#include +#include + +namespace NYdb { +namespace NOdbc { + +SQLSMALLINT GetTypeId(const TType& type); +SQLSMALLINT IsNullable(const TType& type); +SQLULEN GetColumnSize(SQLSMALLINT sqlType); + +std::optional GetDecimalDigits(const TType& type); +std::optional GetRadix(const TType& type); + +} // namespace NOdbc +} // namespace NYdb diff --git a/odbc/src/utils/util.cpp b/odbc/src/utils/util.cpp new file mode 100644 index 00000000000..ba587c6edf1 --- /dev/null +++ b/odbc/src/utils/util.cpp @@ -0,0 +1,170 @@ +#include "util.h" + +#include +#include + +namespace NYdb::NOdbc { + +namespace { + +void TrimInPlace(std::string& value) { + while (!value.empty() && std::isspace(static_cast(value.front()))) { + value.erase(value.begin()); + } + while (!value.empty() && std::isspace(static_cast(value.back()))) { + value.pop_back(); + } +} + +} // namespace + +std::string GetString(SQLCHAR* str, SQLINTEGER length) { + if (!str) { + return {}; + } + if (length == SQL_NTS) { + return std::string(reinterpret_cast(str)); + } + if (length <= 0) { + return {}; + } + size_t size = static_cast(length); + if (str[size - 1] == 0) { + --size; + } + return std::string(reinterpret_cast(str), size); +} + +std::string GetString(SQLWCHAR* str, SQLINTEGER length) { + if (!str) { + return {}; + } + + size_t size = 0; + if (length == SQL_NTS) { + while (str[size] != 0) { + ++size; + } + } else if (length > 0) { + size = static_cast(length); + if (str[size - 1] == 0) { + --size; + } + } else { + return {}; + } + + std::string result; + result.reserve(size); + for (size_t i = 0; i < size; ++i) { + uint32_t codePoint = str[i]; + if (codePoint >= 0xd800 && codePoint <= 0xdbff) { + if (i + 1 < size && str[i + 1] >= 0xdc00 && str[i + 1] <= 0xdfff) { + codePoint = 0x10000 + ((codePoint - 0xd800) << 10) + (str[++i] - 0xdc00); + } else { + codePoint = 0xfffd; + } + } else if (codePoint >= 0xdc00 && codePoint <= 0xdfff) { + codePoint = 0xfffd; + } + + if (codePoint <= 0x7f) { + result.push_back(static_cast(codePoint)); + } else if (codePoint <= 0x7ff) { + result.push_back(static_cast(0xc0 | (codePoint >> 6))); + result.push_back(static_cast(0x80 | (codePoint & 0x3f))); + } else if (codePoint <= 0xffff) { + result.push_back(static_cast(0xe0 | (codePoint >> 12))); + result.push_back(static_cast(0x80 | ((codePoint >> 6) & 0x3f))); + result.push_back(static_cast(0x80 | (codePoint & 0x3f))); + } else { + result.push_back(static_cast(0xf0 | (codePoint >> 18))); + result.push_back(static_cast(0x80 | ((codePoint >> 12) & 0x3f))); + result.push_back(static_cast(0x80 | ((codePoint >> 6) & 0x3f))); + result.push_back(static_cast(0x80 | (codePoint & 0x3f))); + } + } + return result; +} + +bool StartsWithPrefix(const char* s, size_t sLen, const char* prefix, size_t prefixLen) { + if (sLen < prefixLen) { + return false; + } + for (size_t i = 0; i < prefixLen; ++i) { + if (std::tolower(static_cast(s[i])) != + std::tolower(static_cast(prefix[i]))) { + return false; + } + } + return true; +} + +TConnectionStringEntries ParseConnectionStringEntries(std::string_view connectionString) { + TConnectionStringEntries entries; + size_t pos = 0; + while (pos < connectionString.size()) { + const size_t eq = connectionString.find('=', pos); + if (eq == std::string::npos) { + break; + } + std::string key(connectionString.substr(pos, eq - pos)); + TrimInPlace(key); + if (key.empty()) { + break; + } + + size_t valueStart = eq + 1; + size_t valueEnd = connectionString.size(); + if (valueStart < connectionString.size() && connectionString[valueStart] == '{') { + ++valueStart; + size_t braceDepth = 1; + size_t i = valueStart; + while (i < connectionString.size() && braceDepth > 0) { + if (connectionString[i] == '{') { + ++braceDepth; + } else if (connectionString[i] == '}') { + --braceDepth; + if (braceDepth == 0) { + valueEnd = i; + pos = i + 1; + if (pos < connectionString.size() && connectionString[pos] == ';') { + ++pos; + } + break; + } + } + ++i; + } + if (braceDepth != 0) { + valueEnd = connectionString.size(); + pos = connectionString.size(); + } + entries.emplace_back( + std::move(key), std::string(connectionString.substr(valueStart, valueEnd - valueStart))); + continue; + } + + const size_t sc = connectionString.find(';', valueStart); + if (sc != std::string::npos) { + valueEnd = sc; + pos = sc + 1; + } else { + pos = connectionString.size(); + } + std::string val(connectionString.substr(valueStart, valueEnd - valueStart)); + TrimInPlace(val); + entries.emplace_back(std::move(key), std::move(val)); + } + return entries; +} + +std::map ParseConnectionString(std::string_view connectionString) { + std::map params; + for (auto&& [key, value] : ParseConnectionStringEntries(connectionString)) { + params[std::move(key)] = std::move(value); + } + return params; +} + +} // namespace NYdb::NOdbc diff --git a/odbc/src/utils/util.h b/odbc/src/utils/util.h new file mode 100644 index 00000000000..adb5d5d4906 --- /dev/null +++ b/odbc/src/utils/util.h @@ -0,0 +1,28 @@ +#pragma once + +#include + +#include +#include + +#include +#include +#include +#include +#include + +namespace NYdb::NOdbc { + +std::string GetString(SQLCHAR* str, SQLINTEGER length); + +std::string GetString(SQLWCHAR* str, SQLINTEGER length); + +bool StartsWithPrefix(const char* s, size_t sLen, const char* prefix, size_t prefixLen); + +using TConnectionStringEntries = std::vector>; + +TConnectionStringEntries ParseConnectionStringEntries(std::string_view connectionString); + +std::map ParseConnectionString(std::string_view connectionString); + +} // namespace NYdb::NOdbc diff --git a/odbc/tests/CMakeLists.txt b/odbc/tests/CMakeLists.txt new file mode 100644 index 00000000000..916fa4a0b8d --- /dev/null +++ b/odbc/tests/CMakeLists.txt @@ -0,0 +1,22 @@ +set(YDB_ODBC_TEST_CONFIG_DIR "${CMAKE_BINARY_DIR}/odbc") +file(MAKE_DIRECTORY "${YDB_ODBC_TEST_CONFIG_DIR}") + +set(YDB_ODBC_DSN_SERVER "localhost:2136" CACHE STRING + "YDB endpoint in odbc.ini generated for ODBC integration tests") +set(YDB_ODBC_DSN_DATABASE "/local" CACHE STRING + "YDB database path in odbc.ini generated for ODBC integration tests") + +file(WRITE "${YDB_ODBC_TEST_CONFIG_DIR}/odbc.ini" +"[ODBC Data Sources] +YDB=YDB ODBC Driver + +[YDB] +Driver=YDB +Description=YDB Database Connection +Server=${YDB_ODBC_DSN_SERVER} +Database=${YDB_ODBC_DSN_DATABASE} +AuthMode=Anonymous +") + +add_subdirectory(integration) +add_subdirectory(unit) diff --git a/odbc/tests/integration/CMakeLists.txt b/odbc/tests/integration/CMakeLists.txt new file mode 100644 index 00000000000..a2419c0e459 --- /dev/null +++ b/odbc/tests/integration/CMakeLists.txt @@ -0,0 +1,53 @@ +add_odbc_test(NAME odbc-basic_it + SOURCES + basic_it.cpp +) + +add_odbc_test(NAME odbc-environment_api_it + SOURCES + environment_api_it.cpp +) + +add_odbc_test(NAME odbc-connection_api_it + SOURCES + connection_api_it.cpp +) + +add_odbc_test(NAME odbc-authentication_it + SOURCES + authentication_it.cpp + LINK_LIBRARIES + tests-iam-mocks + client-oauth2-ut-helpers + cpp-testing-unittest +) + +add_odbc_test(NAME odbc-statement_api_it + SOURCES + statement_api_it.cpp +) + +add_odbc_test(NAME odbc-transaction_api_it + SOURCES + transaction_api_it.cpp +) + +add_odbc_test(NAME odbc-error_handling_it + SOURCES + error_handling_it.cpp +) + +add_odbc_test(NAME odbc-metadata_api_it + SOURCES + metadata_api_it.cpp +) + +add_odbc_test(NAME odbc-core_api_it + SOURCES + core_api_it.cpp +) + +add_odbc_test(NAME odbc-descriptor_api_it + SOURCES + descriptor_api_it.cpp +) diff --git a/odbc/tests/integration/authentication_it.cpp b/odbc/tests/integration/authentication_it.cpp new file mode 100644 index 00000000000..4fc9eae283a --- /dev/null +++ b/odbc/tests/integration/authentication_it.cpp @@ -0,0 +1,207 @@ +#include "test_utils.h" + +#include +#include +#include +#include +#include +#include + +#include +#include + +#include +#include +#include +#include + +using namespace NYdb::NTest; + +namespace { + +constexpr std::string_view RootToken = "root@builtin"; + +class TScopedEnvironmentVariable { +public: + TScopedEnvironmentVariable(std::string_view name, std::string_view value) + : Name_(name) + { + if (const char* oldValue = std::getenv(Name_.c_str())) { + OldValue_ = oldValue; + } + setenv(Name_.c_str(), std::string(value).c_str(), 1); + } + + ~TScopedEnvironmentVariable() { + if (OldValue_) { + setenv(Name_.c_str(), OldValue_->c_str(), 1); + } else { + unsetenv(Name_.c_str()); + } + } + +private: + std::string Name_; + std::optional OldValue_; +}; + +class OdbcAuthentication : public ::testing::Test { +protected: + void SetUp() override { + const char* endpoint = std::getenv("YDB_ENDPOINT"); + const char* database = std::getenv("YDB_DATABASE"); + if (!endpoint || !database) { + GTEST_SKIP() << "Authentication integration tests require the IAM-enabled YDB fixture"; + } + Endpoint_ = endpoint; + Database_ = database; + AllocEnv(&Env_); + } + + void TearDown() override { + Disconnect(); + if (Env_ != SQL_NULL_HENV) { + SQLFreeHandle(SQL_HANDLE_ENV, Env_); + } + } + + void Connect(std::string_view authenticationAttributes) { + Disconnect(); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, Env_, &Dbc_), SQL_SUCCESS); + std::string connectionString = "Driver=" ODBC_DRIVER_PATH ";Endpoint=" + Endpoint_ + + ";Database=" + Database_ + ";" + std::string(authenticationAttributes); + const SQLRETURN rc = SQLDriverConnect( + Dbc_, nullptr, reinterpret_cast(connectionString.data()), SQL_NTS, + nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT); + CHECK_ODBC_OK(rc, Dbc_, SQL_HANDLE_DBC); + } + + void Disconnect() { + if (Dbc_ != SQL_NULL_HDBC) { + SQLDisconnect(Dbc_); + SQLFreeHandle(SQL_HANDLE_DBC, Dbc_); + Dbc_ = SQL_NULL_HDBC; + } + } + + void Execute(std::string_view query) { + SQLHSTMT statement = SQL_NULL_HSTMT; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, Dbc_, &statement), SQL_SUCCESS); + std::string queryString(query); + const SQLRETURN rc = SQLExecDirect( + statement, reinterpret_cast(queryString.data()), SQL_NTS); + CHECK_ODBC_OK(rc, statement, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, statement); + } + + void SelectOne() { + Execute("SELECT 1"); + } + + void ExpectSelectOneAuthFailure() { + SQLHSTMT statement = SQL_NULL_HSTMT; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, Dbc_, &statement), SQL_SUCCESS); + SQLCHAR query[] = "SELECT 1"; + const SQLRETURN rc = SQLExecDirect(statement, query, SQL_NTS); + EXPECT_EQ(rc, SQL_ERROR); + if (rc == SQL_ERROR) { + const std::string error = GetOdbcError(statement, SQL_HANDLE_STMT); + EXPECT_TRUE(SqlStatePrefix(error, "28000")) << error; + } + EXPECT_EQ(SQLFreeHandle(SQL_HANDLE_STMT, statement), SQL_SUCCESS); + } + + SQLHENV Env_ = SQL_NULL_HENV; + SQLHDBC Dbc_ = SQL_NULL_HDBC; + std::string Endpoint_; + std::string Database_; +}; + +} // namespace + +TEST_F(OdbcAuthentication, TokenAndAccessTokenAlias) { + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=Token;Token=root@builtin;")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); + + ASSERT_NO_FATAL_FAILURE(Connect("AccessToken=root@builtin;")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); +} + +TEST_F(OdbcAuthentication, Anonymous) { + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=Anonymous;")); + ASSERT_NO_FATAL_FAILURE(ExpectSelectOneAuthFailure()); +} + +TEST_F(OdbcAuthentication, StaticUserAndPasswordAliases) { + const char* user = std::getenv("YDB_ODBC_STATIC_USER"); + const char* password = std::getenv("YDB_ODBC_STATIC_PASSWORD"); + if (!user || !password) { + GTEST_SKIP() << "Static authentication requires credentials provisioned by the IAM fixture"; + } + + ASSERT_NO_FATAL_FAILURE(Connect( + "AuthMode=Static;UID=" + std::string(user) + ";PWD=" + std::string(password) + ";")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); +} + +TEST_F(OdbcAuthentication, MetadataService) { + TMetadataServer server; + server.SetResponse(HTTP_OK, MakeTokenResponse(std::string(RootToken), 3600)); + + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=Metadata;MetadataHost=127.0.0.1;MetadataPort=" + + std::to_string(server.Port) + ";")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); + + EXPECT_GE(server.GetRequestCount(), 1); + AssertMetadataRequestShape(server); +} + +TEST_F(OdbcAuthentication, ServiceAccountFileAndAlias) { + TIamTokenServiceStub stub; + stub.SetResponseToken(std::string(RootToken)); + TIamGrpcServer server(&stub); + ASSERT_TRUE(server.Start()); + + TTempDir tempDirectory; + const TString keyPath = tempDirectory.Path() / "service-account.json"; + TFileOutput(keyPath).Write(MakeJwtKeyFileContent()); + + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=ServiceAccount;SaFile=" + std::string(keyPath) + + ";IamEndpoint=grpc://" + server.Endpoint() + ";")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); + + EXPECT_GE(stub.GetRequestCount(), 1); + ASSERT_TRUE(stub.HasLastRequest()); + AssertIamJwt(stub.GetLastRequest().jwt()); +} + +TEST_F(OdbcAuthentication, OAuth2TokenExchangeFile) { + TTestTokenExchangeServer server; + server.Check.ExpectedInputParams.emplace("grant_type", "urn:ietf:params:oauth:grant-type:token-exchange"); + server.Check.ExpectedInputParams.emplace("requested_token_type", "urn:ietf:params:oauth:token-type:access_token"); + server.Check.ExpectedInputParams.emplace("subject_token", "odbc-subject-token"); + server.Check.ExpectedInputParams.emplace("subject_token_type", "urn:ietf:params:oauth:token-type:access_token"); + server.Check.Response = + R"({"access_token":"root@builtin","token_type":"bearer","expires_in":3600})"; + + TTempDir tempDirectory; + const TString configPath = tempDirectory.Path() / "oauth2.json"; + TFileOutput(configPath).Write( + R"({"subject-credentials":{"type":"Fixed","token":"odbc-subject-token","token-type":"urn:ietf:params:oauth:token-type:access_token"}})"); + + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=OAuth2;OAuth2KeyFile=" + std::string(configPath) + + ";IamEndpoint=" + server.GetEndpoint() + ";")); + // The local IAM fixture accepts builtin tokens, not OAuth "Bearer" credentials. + ASSERT_NO_FATAL_FAILURE(ExpectSelectOneAuthFailure()); + server.CheckExpectations(); +} + +TEST_F(OdbcAuthentication, EnvironmentAccessToken) { + TScopedEnvironmentVariable serviceAccount("YDB_SERVICE_ACCOUNT_KEY_FILE_CREDENTIALS", ""); + TScopedEnvironmentVariable anonymous("YDB_ANONYMOUS_CREDENTIALS", "0"); + TScopedEnvironmentVariable metadata("YDB_METADATA_CREDENTIALS", "0"); + TScopedEnvironmentVariable oauth2("YDB_OAUTH2_KEY_FILE", ""); + TScopedEnvironmentVariable token("YDB_ACCESS_TOKEN_CREDENTIALS", RootToken); + ASSERT_NO_FATAL_FAILURE(Connect("AuthMode=Environment;")); + ASSERT_NO_FATAL_FAILURE(SelectOne()); +} diff --git a/odbc/tests/integration/basic_it.cpp b/odbc/tests/integration/basic_it.cpp new file mode 100644 index 00000000000..e7af877b37a --- /dev/null +++ b/odbc/tests/integration/basic_it.cpp @@ -0,0 +1,126 @@ +#include "test_utils.h" + +TEST(OdbcBasic, SimpleQuery) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + CHECK_ODBC_OK(SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE), dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + // Simple query + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS one, 'abc' AS str", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLINTEGER ival = 0; + char sval[16] = {0}; + SQLLEN ival_ind = 0, sval_ind = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &ival, 0, &ival_ind), SQL_SUCCESS); + ASSERT_EQ(SQLGetData(stmt, 2, SQL_C_CHAR, sval, sizeof(sval), &sval_ind), SQL_SUCCESS); + ASSERT_EQ(ival, 1); + ASSERT_STREQ(sval, "abc"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(OdbcBasic, ParameterizedQuery) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + CHECK_ODBC_OK(SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE), dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLCHAR query[] = R"( + DECLARE $p1 AS Int32?; + SELECT $p1 + 10 AS res; + )"; + + // Parameterized query + CHECK_ODBC_OK(SQLPrepare(stmt, query, SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER param = 5; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, ¶m, 0, nullptr), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLINTEGER res = 0; + SQLLEN res_ind = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &res, 0, &res_ind), SQL_SUCCESS); + ASSERT_EQ(res, 15); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(OdbcBasic, ColumnBinding) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + CHECK_ODBC_OK(SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE), dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLCHAR query_ddl[] = R"( + DROP TABLE IF EXISTS test_bind; + CREATE TABLE test_bind (id Int32, name Text, PRIMARY KEY (id)); + )"; + + SQLCHAR query[] = R"( + UPSERT INTO test_bind (id, name) VALUES (1, 'foo'), (2, 'bar'); + SELECT id, name FROM test_bind ORDER BY id; + )"; + + CHECK_ODBC_OK(SQLExecDirect(stmt, query_ddl, SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLExecDirect(stmt, query, SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLINTEGER id = 0; + char name[16] = {0}; + SQLLEN id_ind = 0, name_ind = 0; + ASSERT_EQ(SQLBindCol(stmt, 1, SQL_C_LONG, &id, 0, &id_ind), SQL_SUCCESS); + ASSERT_EQ(SQLBindCol(stmt, 2, SQL_C_CHAR, name, sizeof(name), &name_ind), SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(id, 1); + ASSERT_STREQ(name, "foo"); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(id, 2); + ASSERT_STREQ(name, "bar"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(OdbcBasic, SQLConnect) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + CHECK_ODBC_OK(SQLConnect(dbc, (SQLCHAR*)"YDB", SQL_NTS, (SQLCHAR*)"", SQL_NTS, (SQLCHAR*)"", SQL_NTS), + dbc, SQL_HANDLE_DBC); + + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLINTEGER val; + SQLLEN ind; + SQLBindCol(stmt, 1, SQL_C_SLONG, &val, 0, &ind); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(val, 1); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/connection_api_it.cpp b/odbc/tests/integration/connection_api_it.cpp new file mode 100644 index 00000000000..4d029aeaa39 --- /dev/null +++ b/odbc/tests/integration/connection_api_it.cpp @@ -0,0 +1,276 @@ +#include "test_utils.h" + +TEST(ConnectionApi, AllocFreeEnvHandle) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_ENV, env), SQL_SUCCESS); +} + +TEST(ConnectionApi, AllocFreeDbcHandle) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_DBC, dbc), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, AllocFreeHandleInvalid) { + SQLHENV env; + SQLRETURN rc = SQLAllocHandle(999, SQL_NULL_HANDLE, &env); + ASSERT_TRUE(rc == SQL_ERROR || rc == SQL_INVALID_HANDLE); +} + +TEST(ConnectionApi, SQLConnectWithDSN) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLConnect(dbc, (SQLCHAR*)"YDB", SQL_NTS, (SQLCHAR*)"", SQL_NTS, (SQLCHAR*)"", SQL_NTS), + dbc, SQL_HANDLE_DBC); + + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectComplete) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + SQLCHAR outStr[256]; + SQLSMALLINT outLen; + SQLRETURN rc = SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, + outStr, sizeof(outStr), &outLen, SQL_DRIVER_COMPLETE); + CHECK_ODBC_OK(rc, dbc, SQL_HANDLE_DBC); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectNoPrompt) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + SQLRETURN rc = SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, + nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT); + CHECK_ODBC_OK(rc, dbc, SQL_HANDLE_DBC); + + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectInvalidConnString) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + SQLRETURN rc = SQLDriverConnect(dbc, nullptr, (SQLCHAR*)"InvalidParam=test", SQL_NTS, + nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT); + ASSERT_EQ(rc, SQL_ERROR); + + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectIgnoresUnrecognizedAttributes) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + SQLCHAR connectionString[] = + "Driver=" ODBC_DRIVER_PATH + ";Endpoint=localhost:2136;Database=/local;APP=PowerBI;WSID=desktop;Timeout=30;"; + const SQLRETURN rc = SQLDriverConnect( + dbc, nullptr, connectionString, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT); + ASSERT_EQ(rc, SQL_SUCCESS_WITH_INFO) << GetOdbcError(dbc, SQL_HANDLE_DBC); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), "01S00")); + + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, reinterpret_cast(const_cast("SELECT 1")), SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectValidatesAuthenticationSettings) { + SQLHENV env; + AllocEnv(&env); + + const struct { + const char* ConnectionString; + const char* SqlState; + } cases[] = { + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;AuthMode=None;", "28000"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;Token=a;UID=b;PWD=c;", "28000"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;AuthMode=Static;UID=b;", "28000"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;AuthMode=Metadata;MetadataPort=70000;", "HY024"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;AuthMode=ServiceAccount;SaFile=/missing/sa.json;", "08001"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;AuthMode=OAuth2;OAuth2KeyFile=/missing/oauth2.json;", "08001"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;ClientCertificate=client.pem;", "08001"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=grpc://localhost:2136;Database=/local;CaFile=ca.pem;", "HY024"}, + {"Driver=" ODBC_DRIVER_PATH ";Endpoint=localhost:2136;Database=/local;RootCertificate=/missing/ca.pem;", "08001"}, + }; + + for (const auto& testCase : cases) { + SQLHDBC dbc; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + const SQLRETURN rc = SQLDriverConnect( + dbc, nullptr, reinterpret_cast(const_cast(testCase.ConnectionString)), SQL_NTS, + nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT); + ASSERT_EQ(rc, SQL_ERROR) << testCase.ConnectionString; + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), testCase.SqlState)) + << testCase.ConnectionString << ": " << GetOdbcError(dbc, SQL_HANDLE_DBC); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + } + + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLDriverConnectSupportsAliasesAndDsnOverlay) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + SQLCHAR connectionString[] = + "DSN=YDB;Endpoint=grpc://127.0.0.1:2136;AuthMode=Token;AccessToken=ignored-by-anonymous-server;"; + CHECK_ODBC_OK(SQLDriverConnect( + dbc, nullptr, connectionString, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_NOPROMPT), + dbc, SQL_HANDLE_DBC); + + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, SQLConnectMissingDSN) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + + SQLRETURN rc = SQLConnect(dbc, (SQLCHAR*)"NONEXISTENT_DSN", SQL_NTS, (SQLCHAR*)"", SQL_NTS, (SQLCHAR*)"", SQL_NTS); + ASSERT_EQ(rc, SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), "IM002")); + + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, EnvAttrOdbcVersion) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + ASSERT_NE(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, nullptr, 0), SQL_SUCCESS); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, EnvAttrOutputNts) { + SQLHENV env; + AllocEnv(&env); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, (void*)SQL_TRUE, 0), SQL_SUCCESS); + ASSERT_NE(SQLSetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, (void*)SQL_FALSE, 0), SQL_SUCCESS); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, ConnAttrAccessMode) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + SQLUINTEGER mode; + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_ACCESS_MODE, &mode, sizeof(mode), nullptr), SQL_SUCCESS); + ASSERT_EQ(mode, SQL_MODE_READ_WRITE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_ACCESS_MODE, (SQLPOINTER)SQL_MODE_READ_ONLY, 0), + dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_ACCESS_MODE, &mode, sizeof(mode), nullptr), SQL_SUCCESS); + ASSERT_EQ(mode, SQL_MODE_READ_ONLY); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_ACCESS_MODE, (SQLPOINTER)SQL_MODE_READ_WRITE, 0), + dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLSetConnectAttr(dbc, SQL_ATTR_ACCESS_MODE, (SQLPOINTER)9999, 0), SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), "HY024")); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, ConnAttrCurrentCatalog) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + char catalog[256]; + SQLINTEGER len; + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_CURRENT_CATALOG, catalog, sizeof(catalog), &len), SQL_SUCCESS); + ASSERT_STREQ(catalog, "/local"); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_CURRENT_CATALOG, (SQLPOINTER)"/local/test", SQL_NTS), + dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_CURRENT_CATALOG, catalog, sizeof(catalog), &len), SQL_SUCCESS); + ASSERT_STREQ(catalog, "/local/test"); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ConnectionApi, ConnAttrCurrentCatalogAffectsQueries) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS `/local/cat_a/probe`", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE `/local/cat_a/probe` (id Int32, value Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO `/local/cat_a/probe` (id, value) VALUES (1, 100)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS `/local/cat_b/probe`", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE `/local/cat_b/probe` (id Int32, value Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO `/local/cat_b/probe` (id, value) VALUES (1, 200)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_CURRENT_CATALOG, (SQLPOINTER)"/local/cat_a", SQL_NTS), + dbc, SQL_HANDLE_DBC); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT value FROM probe WHERE id = 1", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER value; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &value, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(value, 100); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_CURRENT_CATALOG, (SQLPOINTER)"/local/cat_b", SQL_NTS), + dbc, SQL_HANDLE_DBC); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT value FROM probe WHERE id = 1", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &value, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(value, 200); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/core_api_it.cpp b/odbc/tests/integration/core_api_it.cpp new file mode 100644 index 00000000000..3a2a263a01d --- /dev/null +++ b/odbc/tests/integration/core_api_it.cpp @@ -0,0 +1,397 @@ +#include "test_utils.h" + +#include + +#ifndef SQL_ODBC_INTERFACE_CONFORMANCE +#define SQL_ODBC_INTERFACE_CONFORMANCE 169 +#endif + +TEST(CoreApi, SQLGetTypeInfoAll) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLGetTypeInfo(stmt, SQL_ALL_TYPES), stmt, SQL_HANDLE_STMT); + char typeName[64] = {}; + SQLLEN indicator = 0; + SQLBindCol(stmt, 1, SQL_C_CHAR, typeName, sizeof(typeName), &indicator); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + EXPECT_TRUE(std::strstr(typeName, "bigint") != nullptr); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLGetTypeInfoFilter) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLGetTypeInfo(stmt, SQL_INTEGER), stmt, SQL_HANDLE_STMT); + SQLINTEGER dataType = 0; + SQLLEN indicator = 0; + SQLBindCol(stmt, 2, SQL_C_LONG, &dataType, 0, &indicator); + int rowCount = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ASSERT_EQ(dataType, SQL_INTEGER); + ++rowCount; + } + ASSERT_GT(rowCount, 0); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLNumParamsQuestionMarks) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT ? + ?", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLSMALLINT paramCount = 0; + CHECK_ODBC_OK(SQLNumParams(stmt, ¶mCount), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(paramCount, 2); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLNumParamsDollarParams) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT $p1", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLSMALLINT paramCount = 0; + CHECK_ODBC_OK(SQLNumParams(stmt, ¶mCount), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(paramCount, 1); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLColAttributeName) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS col", SQL_NTS), stmt, SQL_HANDLE_STMT); + char name[64] = {}; + SQLSMALLINT nameLen = 0; + CHECK_ODBC_OK(SQLColAttribute(stmt, 1, SQL_DESC_NAME, name, sizeof(name), &nameLen, nullptr), + stmt, SQL_HANDLE_STMT); + EXPECT_STREQ(name, "col"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLColAttributeType) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS col", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLLEN dataType = 0; + CHECK_ODBC_OK(SQLColAttribute(stmt, 1, SQL_DESC_TYPE, nullptr, 0, nullptr, &dataType), + stmt, SQL_HANDLE_STMT); + EXPECT_EQ(dataType, SQL_INTEGER); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLNativeSqlPassthrough) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + char out[64] = {}; + SQLINTEGER outLen = 0; + CHECK_ODBC_OK(SQLNativeSql(dbc, (SQLCHAR*)"SELECT 1", SQL_NTS, (SQLCHAR*)out, sizeof(out), &outLen), + dbc, SQL_HANDLE_DBC); + EXPECT_STREQ(out, "SELECT 1"); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLSetGetCursorName) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetCursorName(stmt, (SQLCHAR*)"mycursor", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLCHAR name[64] = {}; + SQLSMALLINT nameLen = 0; + CHECK_ODBC_OK(SQLGetCursorName(stmt, name, sizeof(name), &nameLen), stmt, SQL_HANDLE_STMT); + EXPECT_STREQ(reinterpret_cast(name), "mycursor"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLStatisticsEmpty) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLStatistics(stmt, nullptr, 0, nullptr, 0, (SQLCHAR*)"%", SQL_NTS, SQL_INDEX_ALL, SQL_ENSURE), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLSpecialColumnsPrimaryKey) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_special_columns_pk", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_special_columns_pk (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + const char* table = "/local/test_special_columns_pk"; + CHECK_ODBC_OK(SQLSpecialColumns(stmt, SQL_BEST_ROWID, nullptr, 0, nullptr, 0, + (SQLCHAR*)table, SQL_NTS, SQL_SCOPE_SESSION, 0), + stmt, SQL_HANDLE_STMT); + char columnName[64] = {}; + SQLLEN indicator = 0; + SQLBindCol(stmt, 2, SQL_C_CHAR, columnName, sizeof(columnName), &indicator); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + EXPECT_STREQ(columnName, "id"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLGetInfoInterfaceConformance) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + SQLUINTEGER conformance = 0; + SQLSMALLINT outLen = 0; + CHECK_ODBC_OK(SQLGetInfo(dbc, SQL_ODBC_INTERFACE_CONFORMANCE, &conformance, 0, &outLen), + dbc, SQL_HANDLE_DBC); + EXPECT_EQ(conformance, SQL_OIC_CORE); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLForeignKeysEmpty) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLForeignKeys(stmt, nullptr, 0, nullptr, 0, (SQLCHAR*)"%", SQL_NTS, + nullptr, 0, nullptr, 0, (SQLCHAR*)"%", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLSMALLINT colCount = 0; + CHECK_ODBC_OK(SQLNumResultCols(stmt, &colCount), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(colCount, 14); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLPrimaryKeys) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_primary_keys", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_primary_keys (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + const char* table = "/local/test_primary_keys"; + CHECK_ODBC_OK(SQLPrimaryKeys(stmt, nullptr, 0, nullptr, 0, (SQLCHAR*)table, SQL_NTS), + stmt, SQL_HANDLE_STMT); + char columnName[64] = {}; + SQLLEN indicator = 0; + SQLBindCol(stmt, 4, SQL_C_CHAR, columnName, sizeof(columnName), &indicator); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + EXPECT_STREQ(columnName, "id"); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLDescribeParamUnknown) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT ?", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLSMALLINT dataType = 0; + SQLULEN paramSize = 0; + SQLSMALLINT decimalDigits = 0; + SQLSMALLINT nullable = 0; + CHECK_ODBC_OK(SQLDescribeParam(stmt, 1, &dataType, ¶mSize, &decimalDigits, &nullable), + stmt, SQL_HANDLE_STMT); + EXPECT_EQ(dataType, SQL_UNKNOWN_TYPE); + EXPECT_EQ(nullable, SQL_NULLABLE_UNKNOWN); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLDescribeParamBound) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT ?", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER value = 42; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &value, 0, nullptr), + stmt, SQL_HANDLE_STMT); + SQLSMALLINT dataType = 0; + SQLULEN paramSize = 0; + SQLSMALLINT decimalDigits = 0; + SQLSMALLINT nullable = 0; + CHECK_ODBC_OK(SQLDescribeParam(stmt, 1, &dataType, ¶mSize, &decimalDigits, &nullable), + stmt, SQL_HANDLE_STMT); + EXPECT_EQ(dataType, SQL_INTEGER); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLParamDataPutData) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_at_exec", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_at_exec (id Int32, val Text, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"UPSERT INTO test_at_exec (id, val) VALUES (1, ?)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLLEN atExec = SQL_DATA_AT_EXEC; + SQLPOINTER parameterToken = &atExec; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 32, 0, + parameterToken, 0, &atExec), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLExecute(stmt), SQL_NEED_DATA); + SQLPOINTER token = nullptr; + ASSERT_EQ(SQLParamData(stmt, &token), SQL_NEED_DATA); + EXPECT_EQ(token, parameterToken); + const char part1[] = "hel"; + CHECK_ODBC_OK(SQLPutData(stmt, (SQLPOINTER)part1, sizeof(part1) - 1), stmt, SQL_HANDLE_STMT); + const char part2[] = "lo"; + CHECK_ODBC_OK(SQLPutData(stmt, (SQLPOINTER)part2, sizeof(part2) - 1), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLPutData(stmt, nullptr, 0), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLParamData(stmt, &token), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(SQLExecute(stmt), SQL_NEED_DATA); + CHECK_ODBC_OK(SQLCancel(stmt), stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLCancel) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT * FROM AS_TABLE(ListMap(ListFromRange(1u, 1000000u), ($x)->(AsStruct($x AS v))))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLCancel(stmt), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLFreeStmtDrop) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFreeStmt(stmt, SQL_DROP), SQL_SUCCESS); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLParamDataPutDataNts) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_at_exec_nts", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_at_exec_nts (id Int32, val Text, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"UPSERT INTO test_at_exec_nts (id, val) VALUES (2, ?)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLLEN atExec = SQL_DATA_AT_EXEC; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 32, 0, + nullptr, 0, &atExec), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLExecute(stmt), SQL_NEED_DATA); + SQLPOINTER token = nullptr; + ASSERT_EQ(SQLParamData(stmt, &token), SQL_NEED_DATA); + const char payload[] = "nts-value"; + CHECK_ODBC_OK(SQLPutData(stmt, (SQLPOINTER)payload, SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLParamData(stmt, &token), stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(CoreApi, SQLCancelIdle) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLCancel(stmt), stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/descriptor_api_it.cpp b/odbc/tests/integration/descriptor_api_it.cpp new file mode 100644 index 00000000000..02b5c401458 --- /dev/null +++ b/odbc/tests/integration/descriptor_api_it.cpp @@ -0,0 +1,91 @@ +#include "test_utils.h" + +TEST(DescriptorApi, ImplicitImpRowDesc) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS col", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLHDESC ird = SQL_NULL_HDESC; + CHECK_ODBC_OK(SQLGetStmtAttr(stmt, SQL_ATTR_IMP_ROW_DESC, &ird, sizeof(ird), nullptr), + stmt, SQL_HANDLE_STMT); + ASSERT_NE(ird, nullptr); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(DescriptorApi, ImpRowDescCount) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS a, 2 AS b", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLHDESC ird = SQL_NULL_HDESC; + CHECK_ODBC_OK(SQLGetStmtAttr(stmt, SQL_ATTR_IMP_ROW_DESC, &ird, sizeof(ird), nullptr), + stmt, SQL_HANDLE_STMT); + SQLSMALLINT descCount = 0; + CHECK_ODBC_OK(SQLGetDescField(ird, 0, SQL_DESC_COUNT, &descCount, 0, nullptr), ird, SQL_HANDLE_DESC); + SQLSMALLINT numCols = 0; + CHECK_ODBC_OK(SQLNumResultCols(stmt, &numCols), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(descCount, numCols); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(DescriptorApi, AppRowDescBindCol) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 42 AS v", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER value = 0; + SQLLEN indicator = 0; + CHECK_ODBC_OK(SQLBindCol(stmt, 1, SQL_C_LONG, &value, 0, &indicator), stmt, SQL_HANDLE_STMT); + SQLHDESC ard = SQL_NULL_HDESC; + CHECK_ODBC_OK(SQLGetStmtAttr(stmt, SQL_ATTR_APP_ROW_DESC, &ard, sizeof(ard), nullptr), + stmt, SQL_HANDLE_STMT); + SQLSMALLINT type = 0; + SQLSMALLINT subType = 0; + SQLLEN length = 0; + CHECK_ODBC_OK(SQLGetDescRec(ard, 1, nullptr, 0, nullptr, &type, &subType, &length, nullptr, nullptr, nullptr), + ard, SQL_HANDLE_DESC); + EXPECT_EQ(type, SQL_C_LONG); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(DescriptorApi, ExplicitDescAllocCopy) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + SQLHDESC src = SQL_NULL_HDESC; + SQLHDESC dst = SQL_NULL_HDESC; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DESC, dbc, &src), SQL_SUCCESS); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DESC, dbc, &dst), SQL_SUCCESS); + SQLSMALLINT type = SQL_C_LONG; + SQLLEN length = sizeof(SQLINTEGER); + CHECK_ODBC_OK(SQLSetDescRec(src, 1, type, SQL_INTEGER, length, 0, 0, nullptr, nullptr, nullptr), + src, SQL_HANDLE_DESC); + CHECK_ODBC_OK(SQLCopyDesc(src, dst), src, SQL_HANDLE_DESC); + SQLSMALLINT outType = 0; + SQLSMALLINT outSubType = 0; + SQLLEN outLen = 0; + CHECK_ODBC_OK(SQLGetDescRec(dst, 1, nullptr, 0, nullptr, &outType, &outSubType, &outLen, nullptr, nullptr, nullptr), + dst, SQL_HANDLE_DESC); + EXPECT_EQ(outType, SQL_C_LONG); + EXPECT_EQ(outLen, length); + SQLFreeHandle(SQL_HANDLE_DESC, src); + SQLFreeHandle(SQL_HANDLE_DESC, dst); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/environment_api_it.cpp b/odbc/tests/integration/environment_api_it.cpp new file mode 100644 index 00000000000..b3395dc620a --- /dev/null +++ b/odbc/tests/integration/environment_api_it.cpp @@ -0,0 +1,225 @@ +#include "test_utils.h" + +TEST(EnvironmentApi, AllocFreeEnv) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_ENV, env), SQL_SUCCESS); +} + +TEST(EnvironmentApi, AllocEnvInvalidType) { + SQLHENV env; + SQLRETURN rc = SQLAllocHandle(999, SQL_NULL_HANDLE, &env); + ASSERT_TRUE(rc == SQL_ERROR || rc == SQL_INVALID_HANDLE); +} + +TEST(EnvironmentApi, FreeInvalidEnvHandle) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_ENV, env), SQL_SUCCESS); + SQLRETURN rc = SQLFreeHandle(SQL_HANDLE_ENV, env); + ASSERT_TRUE(rc == SQL_SUCCESS || rc == SQL_INVALID_HANDLE || rc == SQL_ERROR); +} + +TEST(EnvironmentApi, DoubleFreeEnv) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_ENV, env), SQL_SUCCESS); + // Second free may return error or success depending on implementation + // but should not crash + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, SetOdbcVersion) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, SetOdbcVersionInvalid) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_NE(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, nullptr, 0), SQL_SUCCESS); + ASSERT_NE(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)9999, 0), SQL_SUCCESS); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, GetOdbcVersion) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + + SQLINTEGER version; + ASSERT_EQ(SQLGetEnvAttr(env, SQL_ATTR_ODBC_VERSION, &version, sizeof(version), nullptr), SQL_SUCCESS); + ASSERT_EQ(version, SQL_OV_ODBC3); + + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, SetOutputNtsTrue) { + SQLHENV env; + AllocEnv(&env); + + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, (void*)SQL_TRUE, 0), SQL_SUCCESS); + + SQLINTEGER outputNts; + ASSERT_EQ(SQLGetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, &outputNts, sizeof(outputNts), nullptr), SQL_SUCCESS); + ASSERT_EQ(outputNts, SQL_TRUE); + + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, SetOutputNtsFalse) { + SQLHENV env; + AllocEnv(&env); + ASSERT_NE(SQLSetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, (void*)SQL_FALSE, 0), SQL_SUCCESS); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, GetOutputNtsDefault) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); + SQLINTEGER outputNts; + ASSERT_EQ(SQLGetEnvAttr(env, SQL_ATTR_OUTPUT_NTS, &outputNts, sizeof(outputNts), nullptr), SQL_SUCCESS); + ASSERT_EQ(outputNts, SQL_TRUE); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, SetInvalidEnvAttr) { + SQLHENV env; + AllocEnv(&env); + ASSERT_EQ(SQLSetEnvAttr(env, 9999, (void*)1, 0), SQL_ERROR); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, GetInvalidEnvAttr) { + SQLHENV env; + AllocEnv(&env); + char buffer[256]; + SQLINTEGER len; + ASSERT_EQ(SQLGetEnvAttr(env, 9999, buffer, sizeof(buffer), &len), SQL_ERROR); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, MultipleConnectionsSequential) { + SQLHENV env; + AllocEnv(&env); + for (int i = 0; i < 3; ++i) { + SQLHDBC dbc; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + CHECK_ODBC_OK(SQLConnect(dbc, (SQLCHAR*)"YDB", SQL_NTS, nullptr, 0, nullptr, 0), dbc, SQL_HANDLE_DBC); + SQLHSTMT stmt; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + char query[32]; + snprintf(query, sizeof(query), "SELECT %d", i + 1); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)query, SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + } + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +namespace { + +void StartManualTx(SQLHDBC dbc, SQLHSTMT* stmt) { + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(*stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), *stmt, SQL_HANDLE_STMT); +} + +} // namespace + +TEST(EnvironmentApi, EndTranCommitOnEnv) { + SQLHENV env; + SQLHDBC dbc1, dbc2; + SQLHSTMT stmt1, stmt2; + + AllocEnvAndConnect(&env, &dbc1); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc2), SQL_SUCCESS); + SQLRETURN rc = SQLDriverConnect( + dbc2, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE); + CHECK_ODBC_OK(rc, dbc2, SQL_HANDLE_DBC); + + StartManualTx(dbc1, &stmt1); + StartManualTx(dbc2, &stmt2); + + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_ENV, env, SQL_COMMIT), env, SQL_HANDLE_ENV); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt1); + SQLFreeHandle(SQL_HANDLE_STMT, stmt2); + SQLDisconnect(dbc1); + SQLDisconnect(dbc2); + SQLFreeHandle(SQL_HANDLE_DBC, dbc1); + SQLFreeHandle(SQL_HANDLE_DBC, dbc2); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, EndTranRollbackOnEnv) { + SQLHENV env; + SQLHDBC dbc1, dbc2; + SQLHSTMT stmt1, stmt2; + + AllocEnvAndConnect(&env, &dbc1); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc2), SQL_SUCCESS); + SQLRETURN rc = SQLDriverConnect( + dbc2, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE); + CHECK_ODBC_OK(rc, dbc2, SQL_HANDLE_DBC); + + StartManualTx(dbc1, &stmt1); + StartManualTx(dbc2, &stmt2); + + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_ENV, env, SQL_ROLLBACK), env, SQL_HANDLE_ENV); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt1); + SQLFreeHandle(SQL_HANDLE_STMT, stmt2); + SQLDisconnect(dbc1); + SQLDisconnect(dbc2); + SQLFreeHandle(SQL_HANDLE_DBC, dbc1); + SQLFreeHandle(SQL_HANDLE_DBC, dbc2); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, EndTranPartialFailureReturnsInfo) { + SQLHENV env; + SQLHDBC dbc1, dbc2; + SQLHSTMT stmt1, stmt2; + + AllocEnvAndConnect(&env, &dbc1); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc2), SQL_SUCCESS); + SQLRETURN rc = SQLDriverConnect( + dbc2, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE); + CHECK_ODBC_OK(rc, dbc2, SQL_HANDLE_DBC); + + StartManualTx(dbc1, &stmt1); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc2, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), dbc2, SQL_HANDLE_DBC); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc2, &stmt2), SQL_SUCCESS); + (void)SQLExecDirect(stmt2, (SQLCHAR*)"SELECT FROM", SQL_NTS); + + rc = SQLEndTran(SQL_HANDLE_ENV, env, SQL_COMMIT); + ASSERT_TRUE(rc == SQL_SUCCESS || rc == SQL_SUCCESS_WITH_INFO || rc == SQL_ERROR) + << GetOdbcError(env, SQL_HANDLE_ENV); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt1); + SQLFreeHandle(SQL_HANDLE_STMT, stmt2); + SQLDisconnect(dbc1); + SQLDisconnect(dbc2); + SQLFreeHandle(SQL_HANDLE_DBC, dbc1); + SQLFreeHandle(SQL_HANDLE_DBC, dbc2); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(EnvironmentApi, GetDiagRecEnv) { + SQLHENV env; + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, &env), SQL_SUCCESS); + (void)SQLSetEnvAttr(env, 9999, (void*)1, 0); + SQLCHAR sqlState[6]; + SQLINTEGER nativeError; + SQLCHAR msg[256]; + SQLSMALLINT msgLen; + (void)SQLGetDiagRec(SQL_HANDLE_ENV, env, 1, sqlState, &nativeError, msg, sizeof(msg), &msgLen); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/error_handling_it.cpp b/odbc/tests/integration/error_handling_it.cpp new file mode 100644 index 00000000000..96bd01d7ea2 --- /dev/null +++ b/odbc/tests/integration/error_handling_it.cpp @@ -0,0 +1,130 @@ +#include "test_utils.h" + +TEST(ErrorHandling, GetDiagRecAfterError) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + SQLRETURN rc = SQLConnect(dbc, (SQLCHAR*)"NONEXISTENT_DSN", SQL_NTS, + (SQLCHAR*)"", SQL_NTS, (SQLCHAR*)"", SQL_NTS); + ASSERT_EQ(rc, SQL_ERROR); + SQLCHAR sqlState[6]; + SQLINTEGER nativeError; + SQLCHAR msg[256]; + SQLSMALLINT msgLen; + SQLRETURN diagRc = SQLGetDiagRec(SQL_HANDLE_DBC, dbc, 1, sqlState, &nativeError, + msg, sizeof(msg), &msgLen); + ASSERT_TRUE(diagRc == SQL_SUCCESS || diagRc == SQL_SUCCESS_WITH_INFO); + const size_t copiedLen = std::strlen(reinterpret_cast(msg)); + ASSERT_EQ(msg[copiedLen], static_cast(0)); + if (diagRc == SQL_SUCCESS_WITH_INFO) { + ASSERT_GE(static_cast(msgLen), copiedLen); + } else { + ASSERT_EQ(static_cast(msgLen), copiedLen); + } + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ErrorHandling, GetDiagRecMultipleErrors) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLExecDirect(stmt, (SQLCHAR*)"INVALID SYNTAX HERE", SQL_NTS); + + SQLSMALLINT numRecs; + SQLGetDiagField(SQL_HANDLE_STMT, stmt, 0, SQL_DIAG_NUMBER, &numRecs, 0, nullptr); + + for (SQLSMALLINT i = 1; i <= numRecs; ++i) { + SQLCHAR sqlState[6]; + SQLINTEGER nativeError; + SQLCHAR msg[256]; + SQLSMALLINT msgLen; + SQLRETURN rc = SQLGetDiagRec(SQL_HANDLE_STMT, stmt, i, sqlState, &nativeError, + msg, sizeof(msg), &msgLen); + ASSERT_TRUE(rc == SQL_SUCCESS || rc == SQL_SUCCESS_WITH_INFO); + } + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ErrorHandling, GetDiagFieldState) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"SELECT invalid_column FROM nonexistent_table", SQL_NTS); + SQLCHAR sqlState[6]; + SQLRETURN rc = SQLGetDiagField(SQL_HANDLE_STMT, stmt, 1, SQL_DIAG_SQLSTATE, + sqlState, sizeof(sqlState), nullptr); + ASSERT_TRUE(rc == SQL_SUCCESS || rc == SQL_SUCCESS_WITH_INFO); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ErrorHandling, GetDiagFieldNativeError) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"SELECT * FROM nonexistent_table", SQL_NTS); + SQLINTEGER nativeError; + SQLRETURN rc = SQLGetDiagField(SQL_HANDLE_STMT, stmt, 1, SQL_DIAG_NATIVE, + &nativeError, sizeof(nativeError), nullptr); + ASSERT_TRUE(rc == SQL_SUCCESS || rc == SQL_SUCCESS_WITH_INFO); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + + +TEST(ErrorHandling, SuccessWithInfo) { + SQLHENV env; + SQLHDBC dbc; + AllocEnv(&env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, env, &dbc), SQL_SUCCESS); + SQLCHAR outStr[10]; + SQLSMALLINT outLen; + SQLRETURN rc = SQLDriverConnect(dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, + outStr, sizeof(outStr), &outLen, SQL_DRIVER_NOPROMPT); + if (rc == SQL_SUCCESS_WITH_INFO) { + SQLCHAR sqlState[6]; + SQLINTEGER nativeError; + SQLCHAR msg[256]; + SQLSMALLINT msgLen; + SQLGetDiagRec(SQL_HANDLE_DBC, dbc, 1, sqlState, &nativeError, msg, sizeof(msg), &msgLen); + } + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(ErrorHandling, ClearErrors) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"SELECT * FROM nonexistent_table", SQL_NTS); + SQLSMALLINT numRecs1; + SQLGetDiagField(SQL_HANDLE_STMT, stmt, 0, SQL_DIAG_NUMBER, &numRecs1, 0, nullptr); + ASSERT_GT(numRecs1, 0); + SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS); + SQLSMALLINT numRecs2; + SQLGetDiagField(SQL_HANDLE_STMT, stmt, 0, SQL_DIAG_NUMBER, &numRecs2, 0, nullptr); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/metadata_api_it.cpp b/odbc/tests/integration/metadata_api_it.cpp new file mode 100644 index 00000000000..e93bb97bb4b --- /dev/null +++ b/odbc/tests/integration/metadata_api_it.cpp @@ -0,0 +1,280 @@ +#include "test_utils.h" + +#ifndef SQL_ATTR_METADATA_ID +#define SQL_ATTR_METADATA_ID 10029 +#endif + +TEST(MetadataApi, SQLTablesAll) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, nullptr, 0, nullptr, 0), + stmt, SQL_HANDLE_STMT); + int rowCount = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++rowCount; + } + ASSERT_GT(rowCount, 0); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLTablesWithPattern) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_metadata_pattern_a", SQL_NTS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_metadata_pattern_b", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_metadata_pattern_a (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_metadata_pattern_b (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)"%/test_metadata_pattern_%", SQL_NTS, nullptr, 0), + stmt, SQL_HANDLE_STMT); + int tableCount = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++tableCount; + } + ASSERT_EQ(tableCount, 2); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLTablesExactMatch) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_exact_table", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_exact_table (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_METADATA_ID, (SQLPOINTER)(uintptr_t)SQL_TRUE, 0), + stmt, SQL_HANDLE_STMT); + const std::string exactPath = "/local/test_exact_table"; + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)exactPath.c_str(), SQL_NTS, nullptr, 0), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLTablesLikePatternWithMetadataId) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_meta_table_1", SQL_NTS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_meta_table_2", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_meta_table_1 (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_meta_table_2 (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + SQLULEN metadataId = SQL_FALSE; + ASSERT_EQ(SQLGetStmtAttr(stmt, SQL_ATTR_METADATA_ID, &metadataId, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(metadataId, SQL_FALSE); + const char* likePattern = "%/test_meta_table_%"; + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)likePattern, SQL_NTS, (SQLCHAR*)"TABLE", SQL_NTS), + stmt, SQL_HANDLE_STMT); + int tableRows = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++tableRows; + } + ASSERT_EQ(tableRows, 2); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_METADATA_ID, (SQLPOINTER)(uintptr_t)SQL_TRUE, 0), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLGetStmtAttr(stmt, SQL_ATTR_METADATA_ID, &metadataId, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(metadataId, SQL_TRUE); + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)likePattern, SQL_NTS, (SQLCHAR*)"TABLE", SQL_NTS), + stmt, SQL_HANDLE_STMT); + tableRows = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++tableRows; + } + ASSERT_EQ(tableRows, 0); + SQLFreeStmt(stmt, SQL_CLOSE); + const std::string exactPath = "/local/test_meta_table_1"; + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)exactPath.c_str(), SQL_NTS, (SQLCHAR*)"TABLE", SQL_NTS), + stmt, SQL_HANDLE_STMT); + tableRows = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++tableRows; + } + ASSERT_EQ(tableRows, 1); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_METADATA_ID, (SQLPOINTER)(uintptr_t)SQL_FALSE, 0), + stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLColumnsAll) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_columns_all", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_columns_all (id Int32, name Text, value Int32, PRIMARY KEY (id))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLColumns(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)"/local/test_columns_all", SQL_NTS, nullptr, 0), + stmt, SQL_HANDLE_STMT); + int colCount = 0; + while (SQLFetch(stmt) == SQL_SUCCESS) { + ++colCount; + } + ASSERT_EQ(colCount, 3); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLColumnsWithPattern) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_columns_pattern", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_columns_pattern (id Int32, value_x Int32, value_y Int32, PRIMARY KEY (id))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + constexpr SQLUSMALLINT kColumnNameCol = 4; + char colName[256] = {}; + SQLLEN colInd = 0; + CHECK_ODBC_OK(SQLColumns(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)"/local/test_columns_pattern", SQL_NTS, + (SQLCHAR*)"val%", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLGetData(stmt, kColumnNameCol, SQL_C_CHAR, colName, sizeof(colName), &colInd), SQL_SUCCESS); + ASSERT_STREQ(colName, "value_x"); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLGetData(stmt, kColumnNameCol, SQL_C_CHAR, colName, sizeof(colName), &colInd), SQL_SUCCESS); + ASSERT_STREQ(colName, "value_y"); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLColumnsMetadataId) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_columns_metadata", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_columns_metadata (id Int32, value_x Int32, PRIMARY KEY (id))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + const std::string exactTable = "/local/test_columns_metadata"; + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_METADATA_ID, (SQLPOINTER)(uintptr_t)SQL_TRUE, 0), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLColumns(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)exactTable.c_str(), SQL_NTS, + (SQLCHAR*)"val%", SQL_NTS), + SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLColumns(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)exactTable.c_str(), SQL_NTS, + (SQLCHAR*)"nonexistent_col%", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_METADATA_ID, (SQLPOINTER)(uintptr_t)SQL_FALSE, 0), + stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, SQLTablesFilterByType) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_type_filter", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_type_filter (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)"/local/test_type_filter", SQL_NTS, + (SQLCHAR*)"VIEW", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLTables(stmt, nullptr, 0, nullptr, 0, + (SQLCHAR*)"/local/test_type_filter", SQL_NTS, + (SQLCHAR*)"TABLE", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(MetadataApi, DdlWithComment) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_ddl_comment", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"/* ddl */ CREATE TABLE test_ddl_comment (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/statement_api_it.cpp b/odbc/tests/integration/statement_api_it.cpp new file mode 100644 index 00000000000..9dbb1b188b1 --- /dev/null +++ b/odbc/tests/integration/statement_api_it.cpp @@ -0,0 +1,803 @@ +#include "test_utils.h" + +#include + +#ifndef SQL_ATTR_METADATA_ID +#define SQL_ATTR_METADATA_ID 10029 +#endif + +TEST(StatementApi, AllocFreeStmtHandle) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFreeHandle(SQL_HANDLE_STMT, stmt), SQL_SUCCESS); + + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, ExecDirectSimple) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1 AS value", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, ExecDirectMultipleColumns) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT 1 AS int_col, 'hello' AS str_col, CAST(3.14 AS Double) AS float_col", + SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, ExecDirectInvalidSyntax) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + ASSERT_EQ(SQLExecDirect(stmt, (SQLCHAR*)"INVALID SYNTAX HERE", SQL_NTS), SQL_ERROR); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, ExecDirectInvalidTable) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + ASSERT_EQ(SQLExecDirect(stmt, (SQLCHAR*)"SELECT * FROM nonexistent_table", SQL_NTS), SQL_ERROR); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, PrepareAndExecuteWithQuestionMarks) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT ? + ? AS result", SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLINTEGER p1 = 10, p2 = 20; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, &p1, 0, nullptr), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLBindParameter(stmt, 2, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, &p2, 0, nullptr), stmt, SQL_HANDLE_STMT); + + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLINTEGER result = 0; + SQLLEN resultInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &result, 0, &resultInd), SQL_SUCCESS); + ASSERT_EQ(result, 30); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, PrepareAndExecute) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT $p1 + $p2 AS result", SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLINTEGER p1 = 10, p2 = 20; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, &p1, 0, nullptr), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLBindParameter(stmt, 2, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, &p2, 0, nullptr), stmt, SQL_HANDLE_STMT); + + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, PrepareAndExecuteReused) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLPrepare(stmt, (SQLCHAR*)"SELECT $p1", SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER param; + SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, ¶m, 0, nullptr); + param = 100; + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER result; + SQLGetData(stmt, 1, SQL_C_LONG, &result, 0, nullptr); + ASSERT_EQ(result, 100); + SQLCloseCursor(stmt); + param = 200; + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLGetData(stmt, 1, SQL_C_LONG, &result, 0, nullptr); + ASSERT_EQ(result, 200); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, FetchSingleRow) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 42", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, FetchMultipleRows) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT * FROM AS_TABLE(ListMap(ListFromRange(1, 4), ($x) -> (AsStruct($x AS a)))) ORDER BY a", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER value; + SQLLEN ind; + SQLBindCol(stmt, 1, SQL_C_LONG, &value, 0, &ind); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(value, 1); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(value, 2); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(value, 3); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, BindColMultipleTypes) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 42 AS col1, 'test' AS col2", SQL_NTS), + stmt, SQL_HANDLE_STMT); + + SQLINTEGER col1; + char col2[64]; + SQLLEN col1Ind, col2Ind; + + SQLBindCol(stmt, 1, SQL_C_LONG, &col1, 0, &col1Ind); + SQLBindCol(stmt, 2, SQL_C_CHAR, col2, sizeof(col2), &col2Ind); + + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(col1, 42); + ASSERT_STREQ(col2, "test"); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, BindColThenGetData) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 100", SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLINTEGER value; + SQLLEN ind; + SQLBindCol(stmt, 1, SQL_C_LONG, &value, 0, &ind); + + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(value, 100); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, GetDataWithoutBindCol) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 100", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLINTEGER value; + SQLLEN ind; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &value, 0, &ind), SQL_SUCCESS); + ASSERT_EQ(value, 100); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, GetDataMultipleColumns) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1, 'hello world'", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + SQLINTEGER col1; + SQLLEN col1Ind; + SQLGetData(stmt, 1, SQL_C_LONG, &col1, 0, &col1Ind); + ASSERT_EQ(col1, 1); + + char col2[64]; + SQLLEN col2Ind; + SQLGetData(stmt, 2, SQL_C_CHAR, col2, sizeof(col2), &col2Ind); + ASSERT_STREQ(col2, "hello world"); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, CloseCursor) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLCloseCursor(stmt), stmt, SQL_HANDLE_STMT); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, FreeStmtClose) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1", SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLFreeStmt(stmt, SQL_CLOSE), stmt, SQL_HANDLE_STMT); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, FreeStmtResetParams) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLINTEGER param = 42; + SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, ¶m, 0, nullptr); + + CHECK_ODBC_OK(SQLFreeStmt(stmt, SQL_RESET_PARAMS), stmt, SQL_HANDLE_STMT); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, NumResultCols) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT 1, 2, 3, 4, 5", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLSMALLINT numCols; + CHECK_ODBC_OK(SQLNumResultCols(stmt, &numCols), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(numCols, 5); + SQLFetch(stmt); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, RowCount) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS row_count_test", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE row_count_test (id Int32, value Int32, PRIMARY KEY (id))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + + SQLLEN rowCount = -2; + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, -1); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"UPSERT INTO row_count_test (id, value) VALUES (1, 10), (2, 20), (3, 30)", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLLEN diagRowCount = -2; + CHECK_ODBC_OK(SQLGetDiagField(SQL_HANDLE_STMT, stmt, 0, SQL_DIAG_ROW_COUNT, + &diagRowCount, 0, nullptr), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 3); + EXPECT_EQ(diagRowCount, rowCount); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"UPDATE row_count_test SET value = value + 1 WHERE id <= 2", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 2); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"DELETE FROM row_count_test WHERE id = 3", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 1); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"UPDATE row_count_test SET value = 0 WHERE id = 100", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 0); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"PRAGMA TablePathPrefix = \"/local\";\n" + "UPDATE row_count_test SET value = value + 1 WHERE id = 1", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 1); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLPrepare(stmt, + (SQLCHAR*)"DECLARE $p1 AS Int32?;\n" + "UPDATE row_count_test SET value = value + 1 WHERE id = $p1", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER nativeId = 2; + SQLLEN nativeIdLength = 0; + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, &nativeId, 0, &nativeIdLength), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 1); + SQLFreeStmt(stmt, SQL_RESET_PARAMS); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT * FROM row_count_test", + SQL_NTS), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, -1); + + SQLFreeStmt(stmt, SQL_CLOSE); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE row_count_test", SQL_NTS); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, RowCountAggregatesParameterArrays) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS row_count_param_test", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE row_count_param_test (id Int32, value Int32, PRIMARY KEY (id))", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + + CHECK_ODBC_OK(SQLPrepare(stmt, + (SQLCHAR*)"UPSERT INTO row_count_param_test (id, value) VALUES (?, ?)", + SQL_NTS), stmt, SQL_HANDLE_STMT); + SQLINTEGER ids[] = {1, 2, 3}; + SQLINTEGER values[] = {10, 20, 30}; + SQLLEN idLengths[] = {0, 0, 0}; + SQLLEN valueLengths[] = {0, 0, 0}; + SQLUSMALLINT operations[] = {SQL_PARAM_PROCEED, SQL_PARAM_IGNORE, SQL_PARAM_PROCEED}; + SQLUSMALLINT statuses[] = {SQL_PARAM_UNUSED, SQL_PARAM_UNUSED, SQL_PARAM_UNUSED}; + SQLULEN processed = 0; + + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_PARAMSET_SIZE, + reinterpret_cast(3), 0), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_PARAM_OPERATION_PTR, + operations, 0), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_PARAM_STATUS_PTR, + statuses, 0), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_PARAMS_PROCESSED_PTR, + &processed, 0), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLBindParameter(stmt, 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, ids, 0, idLengths), stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLBindParameter(stmt, 2, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, values, 0, valueLengths), stmt, SQL_HANDLE_STMT); + + CHECK_ODBC_OK(SQLExecute(stmt), stmt, SQL_HANDLE_STMT); + SQLLEN rowCount = -1; + CHECK_ODBC_OK(SQLRowCount(stmt, &rowCount), stmt, SQL_HANDLE_STMT); + EXPECT_EQ(rowCount, 2); + EXPECT_EQ(processed, 3); + EXPECT_EQ(statuses[0], SQL_PARAM_SUCCESS); + EXPECT_EQ(statuses[1], SQL_PARAM_UNUSED); + EXPECT_EQ(statuses[2], SQL_PARAM_SUCCESS); + + SQLFreeStmt(stmt, SQL_CLOSE); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE row_count_param_test", SQL_NTS); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, AttrQueryTimeout) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + + SQLUINTEGER timeoutSec = 1; + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_QUERY_TIMEOUT, (SQLPOINTER)(uintptr_t)timeoutSec, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR longQuery[] = + "SELECT COUNT(*) FROM AS_TABLE(ListMap(ListFromRange(1u, 100000000u), ($x)->(AsStruct($x AS v))))"; + ASSERT_EQ(SQLExecDirect(stmt, longQuery, SQL_NTS), SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(stmt, SQL_HANDLE_STMT), "HYT00")); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, AttrMaxRows) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_max_rows", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE test_max_rows (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO test_max_rows (id) VALUES (1), (2)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + const SQLULEN maxRows = 1; + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_MAX_ROWS, (SQLPOINTER)(uintptr_t)maxRows, 0), + stmt, SQL_HANDLE_STMT); + SQLULEN maxRowsOut; + ASSERT_EQ(SQLGetStmtAttr(stmt, SQL_ATTR_MAX_ROWS, &maxRowsOut, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(maxRowsOut, maxRows); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT id FROM test_max_rows ORDER BY id", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, AttrNoScan) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLCHAR selectEscapeFnQuery[] = "SELECT {fn ABS(-12)} AS value"; + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLExecDirect(stmt, selectEscapeFnQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER valueInt = 0; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &valueInt, 0, &valueInd), SQL_SUCCESS); + ASSERT_EQ(valueInt, 12); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_ON, 0), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLExecDirect(stmt, selectEscapeFnQuery, SQL_NTS), SQL_ERROR); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceConvert) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR convertQuery[] = "SELECT {fn CONVERT(42, SQL_SMALLINT)} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, convertQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLSMALLINT valueSmall = 0; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_SSHORT, &valueSmall, 0, &valueInd), SQL_SUCCESS); + ASSERT_EQ(valueSmall, 42); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceDouble) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR convertDoubleQuery[] = "SELECT {fn CONVERT(2.5, SQL_DOUBLE)} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, convertDoubleQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + double valueDouble = 0; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_DOUBLE, &valueDouble, 0, &valueInd), SQL_SUCCESS); + ASSERT_LT(std::fabs(valueDouble - 2.5), 1e-9); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceNested) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR nestedFnQuery[] = "SELECT {fn {fn ABS(-10)}} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, nestedFnQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER valueInt = 0; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &valueInt, 0, &valueInd), SQL_SUCCESS); + ASSERT_EQ(valueInt, 10); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceString) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR asciiLowerQuery[] = "SELECT {fn String::AsciiToLower('AbC')} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, asciiLowerQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + char buf[32] = {}; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_CHAR, buf, sizeof(buf), &valueInd), SQL_SUCCESS); + ASSERT_STREQ(buf, "abc"); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceDate) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR dateQuery[] = "SELECT {d '2024-06-15'} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, dateQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + char buf[32] = {}; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_CHAR, buf, sizeof(buf), &valueInd), SQL_SUCCESS); + ASSERT_STREQ(buf, "2024-06-15"); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, EscapeSequenceTimestamp) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLSetStmtAttr(stmt, SQL_ATTR_NOSCAN, (SQLPOINTER)SQL_NOSCAN_OFF, 0), + stmt, SQL_HANDLE_STMT); + + SQLCHAR tsQuery[] = "SELECT {ts '2024-06-15 14:30:00'} AS value"; + CHECK_ODBC_OK(SQLExecDirect(stmt, tsQuery, SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + char buf[64] = {}; + SQLLEN valueInd = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_CHAR, buf, sizeof(buf), &valueInd), SQL_SUCCESS); + ASSERT_STREQ(buf, "2024-06-15 14:30:00"); + + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, NumericOutOfRange) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT CAST(3000000000 AS Uint64) AS v", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER value = 0; + SQLLEN indicator = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &value, sizeof(value), &indicator), SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(stmt, SQL_HANDLE_STMT), "22003")); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, UpsertAutocommitPersist) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_upsert_persist", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"CREATE TABLE test_upsert_persist (id Int32, val Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"UPSERT INTO test_upsert_persist (id, val) VALUES (1, 42)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, + (SQLCHAR*)"SELECT val FROM test_upsert_persist WHERE id = 1", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER val = 0; + SQLLEN ind = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &val, sizeof(val), &ind), SQL_SUCCESS); + ASSERT_EQ(val, 42); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(StatementApi, SqlCBit) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT true AS b", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + char bitVal = 0; + SQLLEN ind = 0; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_BIT, &bitVal, sizeof(bitVal), &ind), SQL_SUCCESS); + ASSERT_EQ(bitVal, 1); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT false AS b", SQL_NTS), stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_BIT, &bitVal, sizeof(bitVal), &ind), SQL_SUCCESS); + ASSERT_EQ(bitVal, 0); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/integration/test_utils.h b/odbc/tests/integration/test_utils.h new file mode 100644 index 00000000000..b0eae273349 --- /dev/null +++ b/odbc/tests/integration/test_utils.h @@ -0,0 +1,124 @@ +#pragma once + +#include + +#include +#include + +#include +#include +#include +#include + +inline std::string GetOdbcError(SQLHANDLE handle, SQLSMALLINT type) { + SQLCHAR sqlState[6] = {0}; + SQLCHAR message[256] = {0}; + SQLINTEGER nativeError = 0; + SQLSMALLINT textLength = 0; + SQLRETURN rc = SQLGetDiagRec(type, handle, 1, sqlState, &nativeError, message, sizeof(message), &textLength); + if (rc == SQL_SUCCESS || rc == SQL_SUCCESS_WITH_INFO) { + return std::string((char*)sqlState) + ": " + (char*)message; + } + return "Unknown ODBC error"; +} + +#define CHECK_ODBC_OK(rc, handle, type) \ + ASSERT_TRUE((rc) == SQL_SUCCESS || (rc) == SQL_SUCCESS_WITH_INFO) << GetOdbcError(handle, type) + +inline const char* kConnStr = "Driver=" ODBC_DRIVER_PATH ";Server=localhost:2136;Database=/local;"; + +inline bool SqlStatePrefix(std::string_view diag, std::string_view state) { + return diag.starts_with(state); +} + +inline void AllocEnv(SQLHENV* env) { + static bool configured = false; + if (!configured) { + setenv("ODBCINI", ODBC_TEST_ODBCINI, 1); + setenv("ODBCSYSINI", ODBC_TEST_ODBCSYSINI, 1); + configured = true; + } + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_ENV, SQL_NULL_HANDLE, env), SQL_SUCCESS); + ASSERT_EQ(SQLSetEnvAttr(*env, SQL_ATTR_ODBC_VERSION, (void*)SQL_OV_ODBC3, 0), SQL_SUCCESS); +} + +inline void AllocEnvAndConnect(SQLHENV* env, SQLHDBC* dbc) { + AllocEnv(env); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_DBC, *env, dbc), SQL_SUCCESS); + SQLRETURN rc = SQLDriverConnect( + *dbc, nullptr, (SQLCHAR*)kConnStr, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE); + CHECK_ODBC_OK(rc, *dbc, SQL_HANDLE_DBC); +} + +// ============================================================================ +// Type and Parameter Utilities +// ============================================================================ + +// Bind integer parameter and return result +inline SQLRETURN BindIntParam(SQLHSTMT stmt, SQLUSMALLINT paramNum, SQLINTEGER* value) { + return SQLBindParameter(stmt, paramNum, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, value, 0, nullptr); +} + +inline SQLRETURN BindInt64Param(SQLHSTMT stmt, SQLUSMALLINT paramNum, SQLBIGINT* value) { + return SQLBindParameter(stmt, paramNum, SQL_PARAM_INPUT, SQL_C_SBIGINT, SQL_BIGINT, 0, 0, value, 0, nullptr); +} + +inline SQLRETURN BindStringParam(SQLHSTMT stmt, SQLUSMALLINT paramNum, char* value, SQLLEN len) { + SQLLEN indicator = (len >= 0) ? len : SQL_NTS; + return SQLBindParameter(stmt, paramNum, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, + 0, 0, value, (indicator == SQL_NTS) ? 0 : indicator, &indicator); +} + +inline SQLRETURN BindDoubleParam(SQLHSTMT stmt, SQLUSMALLINT paramNum, double* value) { + return SQLBindParameter(stmt, paramNum, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DOUBLE, 0, 0, value, 0, nullptr); +} + +inline SQLRETURN BindNullParam(SQLHSTMT stmt, SQLUSMALLINT paramNum, SQLINTEGER* placeholder) { + static SQLLEN nullIndicator = SQL_NULL_DATA; + return SQLBindParameter(stmt, paramNum, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, + 0, 0, placeholder, 0, &nullIndicator); +} + +// Fetch and verify integer result +inline SQLINTEGER FetchIntResult(SQLHSTMT stmt, SQLUSMALLINT colNum = 1) { + SQLINTEGER result = 0; + SQLLEN indicator = 0; + SQLBindCol(stmt, colNum, SQL_C_LONG, &result, 0, &indicator); + SQLFetch(stmt); + return result; +} + +inline SQLBIGINT FetchInt64Result(SQLHSTMT stmt, SQLUSMALLINT colNum = 1) { + SQLBIGINT result = 0; + SQLLEN indicator = 0; + SQLBindCol(stmt, colNum, SQL_C_SBIGINT, &result, 0, &indicator); + SQLFetch(stmt); + return result; +} + +inline double FetchDoubleResult(SQLHSTMT stmt, SQLUSMALLINT colNum = 1) { + double result = 0.0; + SQLLEN indicator = 0; + SQLBindCol(stmt, colNum, SQL_C_DOUBLE, &result, 0, &indicator); + SQLFetch(stmt); + return result; +} + +inline std::string FetchStringResult(SQLHSTMT stmt, SQLUSMALLINT colNum = 1, size_t maxLen = 256) { + std::string result(maxLen, '\0'); + SQLLEN indicator = 0; + SQLBindCol(stmt, colNum, SQL_C_CHAR, &result[0], maxLen, &indicator); + SQLFetch(stmt); + if (indicator > 0 && indicator != SQL_NULL_DATA) { + result.resize(indicator); + } + return result; +} + +inline bool IsNullResult(SQLHSTMT stmt, SQLUSMALLINT colNum = 1) { + SQLINTEGER dummy; + SQLLEN indicator = 0; + SQLBindCol(stmt, colNum, SQL_C_LONG, &dummy, 0, &indicator); + SQLFetch(stmt); + return indicator == SQL_NULL_DATA; +} diff --git a/odbc/tests/integration/transaction_api_it.cpp b/odbc/tests/integration/transaction_api_it.cpp new file mode 100644 index 00000000000..b26953e4dd0 --- /dev/null +++ b/odbc/tests/integration/transaction_api_it.cpp @@ -0,0 +1,179 @@ +#include "test_utils.h" + +TEST(TransactionApi, AutocommitDefaultOn) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + + SQLUINTEGER autocommit; + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, &autocommit, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(autocommit, SQL_AUTOCOMMIT_ON); + + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, AutocommitOnOffToggle) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), + dbc, SQL_HANDLE_DBC); + SQLUINTEGER autocommit; + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, &autocommit, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(autocommit, SQL_AUTOCOMMIT_OFF); + + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_ON, 0), + dbc, SQL_HANDLE_DBC); + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, &autocommit, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(autocommit, SQL_AUTOCOMMIT_ON); + + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, AutocommitOffRollback) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_rollback", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE test_rollback (id Int32, value Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), + dbc, SQL_HANDLE_DBC); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO test_rollback (id, value) VALUES (1, 100)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_DBC, dbc, SQL_ROLLBACK), dbc, SQL_HANDLE_DBC); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT value FROM test_rollback WHERE id = 1", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_NO_DATA); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, AutocommitOffCommit) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_commit", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE test_commit (id Int32, value Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), + dbc, SQL_HANDLE_DBC); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO test_commit (id, value) VALUES (1, 200)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_DBC, dbc, SQL_COMMIT), dbc, SQL_HANDLE_DBC); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT value FROM test_commit WHERE id = 1", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER value; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &value, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(value, 200); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, MultipleStatementsInManualTransaction) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_multi_stmt", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE test_multi_stmt (id Int32, value Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), + dbc, SQL_HANDLE_DBC); + for (int i = 1; i <= 5; ++i) { + char query[256]; + snprintf(query, sizeof(query), "UPSERT INTO test_multi_stmt (id, value) VALUES (%d, %d)", i, i * 10); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)query, SQL_NTS), stmt, SQL_HANDLE_STMT); + } + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_DBC, dbc, SQL_COMMIT), dbc, SQL_HANDLE_DBC); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"SELECT COUNT(*) FROM test_multi_stmt", SQL_NTS), + stmt, SQL_HANDLE_STMT); + ASSERT_EQ(SQLFetch(stmt), SQL_SUCCESS); + SQLINTEGER count; + ASSERT_EQ(SQLGetData(stmt, 1, SQL_C_LONG, &count, 0, nullptr), SQL_SUCCESS); + ASSERT_EQ(count, 5); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, SQLEndTranOnEnv) { + SQLHENV env; + SQLHDBC dbc; + SQLHSTMT stmt; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLAllocHandle(SQL_HANDLE_STMT, dbc, &stmt), SQL_SUCCESS); + SQLExecDirect(stmt, (SQLCHAR*)"DROP TABLE IF EXISTS test_env_tran", SQL_NTS); + SQLFreeStmt(stmt, SQL_CLOSE); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"CREATE TABLE test_env_tran (id Int32, PRIMARY KEY (id))", SQL_NTS), + stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLSetConnectAttr(dbc, SQL_ATTR_AUTOCOMMIT, (SQLPOINTER)SQL_AUTOCOMMIT_OFF, 0), + dbc, SQL_HANDLE_DBC); + CHECK_ODBC_OK(SQLExecDirect(stmt, (SQLCHAR*)"UPSERT INTO test_env_tran (id) VALUES (1)", SQL_NTS), + stmt, SQL_HANDLE_STMT); + CHECK_ODBC_OK(SQLEndTran(SQL_HANDLE_ENV, env, SQL_COMMIT), env, SQL_HANDLE_ENV); + SQLFreeHandle(SQL_HANDLE_STMT, stmt); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, SQLEndTranInvalid) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLEndTran(SQL_HANDLE_DBC, dbc, 999), SQL_ERROR); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, TxnIsolationDefault) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + SQLUINTEGER isolation; + ASSERT_EQ(SQLGetConnectAttr(dbc, SQL_ATTR_TXN_ISOLATION, &isolation, sizeof(isolation), nullptr), SQL_SUCCESS); + ASSERT_EQ(isolation, SQL_TXN_SERIALIZABLE); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} + +TEST(TransactionApi, TxnIsolationUnsupportedInReadWrite) { + SQLHENV env; + SQLHDBC dbc; + AllocEnvAndConnect(&env, &dbc); + ASSERT_EQ(SQLSetConnectAttr(dbc, SQL_ATTR_TXN_ISOLATION, (SQLPOINTER)SQL_TXN_READ_COMMITTED, 0), SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), "HYC00")); + ASSERT_EQ(SQLSetConnectAttr(dbc, SQL_ATTR_TXN_ISOLATION, (SQLPOINTER)9999, 0), SQL_ERROR); + EXPECT_TRUE(SqlStatePrefix(GetOdbcError(dbc, SQL_HANDLE_DBC), "HY024")); + SQLDisconnect(dbc); + SQLFreeHandle(SQL_HANDLE_DBC, dbc); + SQLFreeHandle(SQL_HANDLE_ENV, env); +} diff --git a/odbc/tests/unit/CMakeLists.txt b/odbc/tests/unit/CMakeLists.txt new file mode 100644 index 00000000000..c2c18c9bada --- /dev/null +++ b/odbc/tests/unit/CMakeLists.txt @@ -0,0 +1,63 @@ +add_ydb_test(NAME odbc-convert_ut GTEST + SOURCES + convert_ut.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/convert.cpp + INCLUDE_DIRS + ${CMAKE_CURRENT_SOURCE_DIR}/../../src + LINK_LIBRARIES + yutil + YDB-CPP-SDK::Params + api-protos + LABELS + unit +) + +add_ydb_test(NAME odbc-escape_ut GTEST + SOURCES + escape_ut.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/escape.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/sql_type_map.cpp + INCLUDE_DIRS + ${CMAKE_CURRENT_SOURCE_DIR}/../../src + LINK_LIBRARIES + yutil + LABELS + unit +) + +add_ydb_test(NAME odbc-param_rewrite_ut GTEST + SOURCES + param_rewrite_ut.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/param_rewrite.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/sql_type_map.cpp + INCLUDE_DIRS + ${CMAKE_CURRENT_SOURCE_DIR}/../../src + LINK_LIBRARIES + yutil + LABELS + unit +) + +add_ydb_test(NAME odbc-conn_string_ut GTEST + SOURCES + conn_string_ut.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/../../src/utils/util.cpp + INCLUDE_DIRS + ${CMAKE_CURRENT_SOURCE_DIR}/../../src + LINK_LIBRARIES + yutil + YDB-CPP-SDK::Params + LABELS + unit +) + +add_ydb_test(NAME odbc-sql_like_ut GTEST + SOURCES + sql_like_ut.cpp + INCLUDE_DIRS + ${CMAKE_CURRENT_SOURCE_DIR}/../../src + LINK_LIBRARIES + yutil + LABELS + unit +) diff --git a/odbc/tests/unit/conn_string_ut.cpp b/odbc/tests/unit/conn_string_ut.cpp new file mode 100644 index 00000000000..4ed8cec3703 --- /dev/null +++ b/odbc/tests/unit/conn_string_ut.cpp @@ -0,0 +1,42 @@ +#include "utils/util.h" + +#include + +TEST(ConnString, ParsesBraceEscapedSemicolons) { + const auto params = NYdb::NOdbc::ParseConnectionString("Database={path;with;semicolons};Server=host"); + ASSERT_EQ(params.at("Database"), "path;with;semicolons"); + ASSERT_EQ(params.at("Server"), "host"); +} + +TEST(ConnString, ParsesSimplePairs) { + const auto params = NYdb::NOdbc::ParseConnectionString("DSN=YDB;Database=/local;Server=grpc://localhost:2136"); + ASSERT_EQ(params.at("DSN"), "YDB"); + ASSERT_EQ(params.at("Database"), "/local"); + ASSERT_EQ(params.at("Server"), "grpc://localhost:2136"); +} + +TEST(ConnString, TrimsWhitespace) { + const auto params = NYdb::NOdbc::ParseConnectionString(" Database = /local ; Server = host "); + ASSERT_EQ(params.at("Database"), "/local"); + ASSERT_EQ(params.at("Server"), "host"); +} + +TEST(OdbcString, ConvertsUtf16ToUtf8) { + SQLWCHAR text[] = {'Y', 'D', 'B', ' ', 0x041f, 0x0440, 0x0438, 0x0432, 0x0435, 0x0442, 0}; + EXPECT_EQ(NYdb::NOdbc::GetString(text, SQL_NTS), "YDB \xd0\x9f\xd1\x80\xd0\xb8\xd0\xb2\xd0\xb5\xd1\x82"); +} + +TEST(OdbcString, ConvertsUtf16SurrogatePairToUtf8) { + SQLWCHAR text[] = {0xd83d, 0xde80, 0}; + EXPECT_EQ(NYdb::NOdbc::GetString(text, SQL_NTS), "\xf0\x9f\x9a\x80"); +} + +TEST(OdbcString, IgnoresAnsiTerminatorIncludedInExplicitLength) { + SQLCHAR text[] = {'/', 'l', 'o', 'c', 'a', 'l', 0}; + EXPECT_EQ(NYdb::NOdbc::GetString(text, 7), "/local"); +} + +TEST(OdbcString, IgnoresUtf16TerminatorIncludedInExplicitLength) { + SQLWCHAR text[] = {'S', 'E', 'L', 'E', 'C', 'T', ' ', '4', '2', 0}; + EXPECT_EQ(NYdb::NOdbc::GetString(text, 10), "SELECT 42"); +} diff --git a/odbc/tests/unit/convert_ut.cpp b/odbc/tests/unit/convert_ut.cpp new file mode 100644 index 00000000000..eb7863ff219 --- /dev/null +++ b/odbc/tests/unit/convert_ut.cpp @@ -0,0 +1,346 @@ +#include "utils/convert.h" +#undef BOOL + +#include + +#include + +#include + +#include + +using namespace NYdb::NOdbc; +using namespace NYdb; + +template +void CheckProto(const T& value, const std::string& expected) { + std::string protoStr; + google::protobuf::TextFormat::PrintToString(value, &protoStr); + ASSERT_EQ(protoStr, expected); +} + +TEST(OdbcConvert, Int64ToYdb) { + SQLBIGINT v = 42; + TBoundParam param{ + 1, // ParamNumber + SQL_PARAM_INPUT, // InputOutputType + SQL_C_SBIGINT, // ValueType + SQL_BIGINT, // ParameterType + 0, 0, // ColumnSize, DecimalDigits + &v, // ParameterValuePtr + sizeof(v), // BufferLength + nullptr // StrLenOrIndPtr + }; + + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: INT64\n }\n}\n"); + CheckProto(value->GetProto(), "int64_value: 42\n"); +} + +TEST(OdbcConvert, UnsignedCToSignedSqlBigint) { + SQLUBIGINT v = 123; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_UBIGINT, SQL_BIGINT, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: INT64\n }\n}\n"); + CheckProto(value->GetProto(), "int64_value: 123\n"); +} + +TEST(OdbcConvert, DoubleToYdb) { + SQLDOUBLE v = 3.14; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DOUBLE, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: DOUBLE\n }\n}\n"); + CheckProto(value->GetProto(), "double_value: 3.14\n"); +} + +TEST(OdbcConvert, StringToYdbUtf8) { + const char* str = "hello"; + SQLLEN len = 5; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, len, nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: UTF8\n }\n}\n"); + CheckProto(value->GetProto(), "text_value: \"hello\"\n"); +} + +TEST(OdbcConvert, StringToYdbBinary) { + const char* str = "bin\x01\x02"; + SQLLEN len = 5; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_BINARY, SQL_BINARY, 0, 0, (SQLPOINTER)str, len, nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: STRING\n }\n}\n"); + CheckProto(value->GetProto(), "bytes_value: \"bin\\001\\002\"\n"); +} + +TEST(OdbcConvert, Int64NullToYdb) { + SQLBIGINT v = 42; + SQLLEN nullInd = SQL_NULL_DATA; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_SBIGINT, SQL_BIGINT, 0, 0, &v, sizeof(v), &nullInd + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: INT64\n }\n}\n"); + CheckProto(value->GetProto(), "null_flag_value: NULL_VALUE\n"); +} + +TEST(OdbcConvert, StringNullToYdb) { + const char* str = "test"; + SQLLEN nullInd = SQL_NULL_DATA; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, 4, &nullInd + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: UTF8\n }\n}\n"); + CheckProto(value->GetProto(), "null_flag_value: NULL_VALUE\n"); +} + +TEST(OdbcConvert, Int32ToYdb) { + SQLINTEGER v = 42; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetType().GetProto(), "optional_type {\n item {\n type_id: INT32\n }\n}\n"); + CheckProto(value->GetProto(), "int32_value: 42\n"); +} + +TEST(OdbcConvert, Int32NegativeToYdb) { + SQLINTEGER v = -999; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "int32_value: -999\n"); +} + +TEST(OdbcConvert, Int32ZeroToYdb) { + SQLINTEGER v = 0; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "int32_value: 0\n"); +} + +TEST(OdbcConvert, Int32MaxToYdb) { + SQLINTEGER v = 2147483647; // INT32_MAX + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "int32_value: 2147483647\n"); +} + +TEST(OdbcConvert, Int32NullToYdb) { + SQLINTEGER v = 42; + SQLLEN nullInd = SQL_NULL_DATA; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &v, sizeof(v), &nullInd + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "null_flag_value: NULL_VALUE\n"); +} + +TEST(OdbcConvert, Int64NegativeToYdb) { + SQLBIGINT v = -123456789012345LL; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_SBIGINT, SQL_BIGINT, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "int64_value: -123456789012345\n"); +} + +TEST(OdbcConvert, Int64ZeroToYdb) { + SQLBIGINT v = 0; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_SBIGINT, SQL_BIGINT, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "int64_value: 0\n"); +} + +TEST(OdbcConvert, DoubleNegativeToYdb) { + SQLDOUBLE v = -2.71828; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DOUBLE, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); +} + +TEST(OdbcConvert, DoubleZeroToYdb) { + SQLDOUBLE v = 0.0; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DOUBLE, 0, 0, &v, sizeof(v), nullptr + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "double_value: 0\n"); +} + +TEST(OdbcConvert, DoubleNullToYdb) { + SQLDOUBLE v = 3.14; + SQLLEN nullInd = SQL_NULL_DATA; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_DOUBLE, SQL_DOUBLE, 0, 0, &v, sizeof(v), &nullInd + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "null_flag_value: NULL_VALUE\n"); +} + +TEST(OdbcConvert, StringEmptyToYdb) { + const char* str = ""; + SQLLEN len = 0; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, len, &len + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "text_value: \"\"\n"); +} + +TEST(OdbcConvert, StringUnicodeToYdb) { + const char* str = "Привет"; + SQLLEN len = SQL_NTS; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, 0, &len + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); +} + +TEST(OdbcConvert, StringWithLengthToYdb) { + const char* str = "hello world"; + SQLLEN len = 5; // Only "hello" + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, len, &len + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "text_value: \"hello\"\n"); +} + +TEST(OdbcConvert, StringNullTerminatedToYdb) { + const char* str = "test"; + SQLLEN len = SQL_NTS; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 0, 0, (SQLPOINTER)str, 0, &len + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "text_value: \"test\"\n"); +} + + +TEST(OdbcConvert, BinaryNullToYdb) { + const char* data = "\x01\x02\x03"; + SQLLEN nullInd = SQL_NULL_DATA; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_BINARY, SQL_BINARY, 0, 0, (SQLPOINTER)data, 3, &nullInd + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "null_flag_value: NULL_VALUE\n"); +} + +TEST(OdbcConvert, BinaryEmptyToYdb) { + const char* data = ""; + SQLLEN len = 0; + TBoundParam param{ + 1, SQL_PARAM_INPUT, SQL_C_BINARY, SQL_BINARY, 0, 0, (SQLPOINTER)data, len, &len + }; + TParamsBuilder paramsBuilder; + ConvertParam(param, paramsBuilder.AddParam("$p1")); + auto params = paramsBuilder.Build(); + auto value = params.GetValue("$p1"); + ASSERT_TRUE(value); + CheckProto(value->GetProto(), "bytes_value: \"\"\n"); +} diff --git a/odbc/tests/unit/escape_ut.cpp b/odbc/tests/unit/escape_ut.cpp new file mode 100644 index 00000000000..60b3e582e69 --- /dev/null +++ b/odbc/tests/unit/escape_ut.cpp @@ -0,0 +1,71 @@ +#include "utils/escape.h" + +#include + +using NYdb::NOdbc::RewriteOdbcEscapes; + +TEST(OdbcEscapeRewrite, FnUnwraps) { + EXPECT_EQ(RewriteOdbcEscapes("SELECT {fn ABS(-12)} AS v"), "SELECT ABS(-12) AS v"); +} + +TEST(OdbcEscapeRewrite, FnCaseInsensitive) { + EXPECT_EQ(RewriteOdbcEscapes("{FN LOWER('A')}"), "LOWER('A')"); +} + +TEST(OdbcEscapeRewrite, OjUnwraps) { + EXPECT_EQ(RewriteOdbcEscapes("{oj LEFT OUTER JOIN t ON a=b}"), "LEFT OUTER JOIN t ON a=b"); +} + +TEST(OdbcEscapeRewrite, DateLiteral) { + EXPECT_EQ(RewriteOdbcEscapes("SELECT {d '2024-01-01'}"), "SELECT CAST('2024-01-01' AS Date)"); +} + +TEST(OdbcEscapeRewrite, TimeLiteral) { + EXPECT_EQ(RewriteOdbcEscapes("{t '14:30:00'}"), "CAST('14:30:00' AS Time)"); +} + +TEST(OdbcEscapeRewrite, TimestampLiteralNormalizesSpaceToT) { + EXPECT_EQ( + RewriteOdbcEscapes("SELECT {ts '2024-06-15 14:30:00'} AS v"), + "SELECT CAST('2024-06-15T14:30:00Z' AS Datetime) AS v"); +} + +TEST(OdbcEscapeRewrite, TimestampLiteralKeepsExistingZ) { + EXPECT_EQ( + RewriteOdbcEscapes("SELECT {ts '2024-06-15T14:30:00Z'} AS v"), + "SELECT CAST('2024-06-15T14:30:00Z' AS Datetime) AS v"); +} + +TEST(OdbcEscapeRewrite, Call) { + EXPECT_EQ(RewriteOdbcEscapes("{call sp_demo(1, 2)}"), "CALL sp_demo(1, 2)"); +} + +TEST(OdbcEscapeRewrite, OutputCallBecomesPlainCall) { + EXPECT_EQ(RewriteOdbcEscapes("{?= call sp(1)}"), "CALL sp(1)"); +} + +TEST(OdbcEscapeRewrite, EscapeClause) { + EXPECT_EQ(RewriteOdbcEscapes("LIKE 'a%' {escape '\\'}"), "LIKE 'a%' ESCAPE '\\'"); +} + +TEST(OdbcEscapeRewrite, ConvertOdbcToYqlCast) { + EXPECT_EQ( + RewriteOdbcEscapes("SELECT {fn CONVERT(42, SQL_SMALLINT)} AS v"), + "SELECT CAST(42 AS Int16) AS v"); +} + +TEST(OdbcEscapeRewrite, ConvertNestedInFn) { + EXPECT_EQ(RewriteOdbcEscapes("{fn CONVERT(x, SQL_INTEGER)}"), "CAST(x AS Int32)"); +} + +TEST(OdbcEscapeRewrite, NestedFnEscapes) { + EXPECT_EQ(RewriteOdbcEscapes("{fn {fn ABS(1)}}"), "ABS(1)"); +} + +TEST(OdbcEscapeRewrite, UnknownBraceLeftUnchanged) { + EXPECT_EQ(RewriteOdbcEscapes("{not_a_keyword 1}"), "{not_a_keyword 1}"); +} + +TEST(OdbcEscapeRewrite, EmptyInput) { + EXPECT_EQ(RewriteOdbcEscapes(""), ""); +} diff --git a/odbc/tests/unit/param_rewrite_ut.cpp b/odbc/tests/unit/param_rewrite_ut.cpp new file mode 100644 index 00000000000..efb979a4ba9 --- /dev/null +++ b/odbc/tests/unit/param_rewrite_ut.cpp @@ -0,0 +1,60 @@ +#include "utils/bindings.h" +#include "utils/param_rewrite.h" + +#include + +using NYdb::NOdbc::RewriteOdbcQuestionMarks; +using NYdb::NOdbc::CountOdbcParams; +using NYdb::NOdbc::TBoundParam; + +namespace { + +TBoundParam IntParam(SQLUSMALLINT n) { + static SQLINTEGER value = 0; + return {n, SQL_PARAM_INPUT, SQL_C_LONG, SQL_INTEGER, 0, 0, &value, 0, nullptr}; +} + +} // namespace + +TEST(OdbcParamRewrite, RewritesQuestionMarks) { + const std::vector params = {IntParam(1), IntParam(2)}; + const auto result = RewriteOdbcQuestionMarks("SELECT ? + ? AS result", params); + ASSERT_TRUE(result.Success); + EXPECT_EQ(result.Sql, + "DECLARE $p1 AS Int32?;\n" + "DECLARE $p2 AS Int32?;\n" + "SELECT $p1 + $p2 AS result"); +} + +TEST(OdbcParamRewrite, SkipsLiteralAndYqlOptionalSyntax) { + const std::vector params = {IntParam(1)}; + EXPECT_EQ(RewriteOdbcQuestionMarks("SELECT '?', ?", params).Sql, + "DECLARE $p1 AS Int32?;\nSELECT '?', $p1"); + EXPECT_EQ(RewriteOdbcQuestionMarks("DECLARE $p1 AS Int32?;\nSELECT $p1", params).Sql, + "DECLARE $p1 AS Int32?;\nSELECT $p1"); + EXPECT_EQ(RewriteOdbcQuestionMarks("SELECT $p1 + 10", params).Sql, + "DECLARE $p1 AS Int32?;\nSELECT $p1 + 10"); +} + +TEST(OdbcParamRewrite, PrependsDeclareForNativeDollarParams) { + const std::vector params = {IntParam(1), IntParam(2)}; + const auto result = RewriteOdbcQuestionMarks("SELECT $p1 + $p2 AS result", params); + ASSERT_TRUE(result.Success); + EXPECT_EQ(result.Sql, + "DECLARE $p1 AS Int32?;\n" + "DECLARE $p2 AS Int32?;\n" + "SELECT $p1 + $p2 AS result"); +} + +TEST(OdbcParamRewrite, RejectsMismatchedBindCount) { + const auto result = RewriteOdbcQuestionMarks("SELECT ? + ?", {IntParam(1)}); + ASSERT_FALSE(result.Success); + EXPECT_EQ(result.SqlState, "07002"); +} + +TEST(OdbcParamRewrite, CountOdbcParams) { + EXPECT_EQ(CountOdbcParams("SELECT ? + ?"), 2); + EXPECT_EQ(CountOdbcParams("SELECT $p1"), 1); + EXPECT_EQ(CountOdbcParams("SELECT $p1 + $p2"), 2); + EXPECT_EQ(CountOdbcParams("SELECT 1"), 0); +} diff --git a/odbc/tests/unit/sql_like_ut.cpp b/odbc/tests/unit/sql_like_ut.cpp new file mode 100644 index 00000000000..e0b8d87ee01 --- /dev/null +++ b/odbc/tests/unit/sql_like_ut.cpp @@ -0,0 +1,28 @@ +#include "utils/sql_like.h" + +#include + +using NYdb::NOdbc::SqlLikeMatch; + +TEST(SqlLikeMatch, PercentMatchesSubstring) { + EXPECT_TRUE(SqlLikeMatch("/local/foo_bar", "%foo%")); + EXPECT_TRUE(SqlLikeMatch("/local/pfx_foo_sfx", "%foo%")); + EXPECT_FALSE(SqlLikeMatch("/local/other", "%foo%")); +} + +TEST(SqlLikeMatch, UnderscoreMatchesSingleChar) { + EXPECT_TRUE(SqlLikeMatch("a_c", "a_c")); + EXPECT_TRUE(SqlLikeMatch("abc", "a_c")); + EXPECT_FALSE(SqlLikeMatch("abbc", "a_c")); +} + +TEST(SqlLikeMatch, EmptyPatternMatchesOnlyEmptyText) { + EXPECT_TRUE(SqlLikeMatch("", "")); + EXPECT_FALSE(SqlLikeMatch("anything", "")); +} + +TEST(SqlLikeMatch, PercentAtEnds) { + EXPECT_TRUE(SqlLikeMatch("hello", "%hello%")); + EXPECT_TRUE(SqlLikeMatch("hello", "hel%")); + EXPECT_TRUE(SqlLikeMatch("hello", "%llo")); +} diff --git a/scripts/build_cpack_deb_packages.sh b/scripts/build_cpack_deb_packages.sh index 91f0b8941e2..b5663cedb43 100755 --- a/scripts/build_cpack_deb_packages.sh +++ b/scripts/build_cpack_deb_packages.sh @@ -23,6 +23,8 @@ if [ "${YDB_DEB_INSTALL_DEPS:-1}" = "1" ]; then build-essential \ ccache \ cmake \ + dpkg-dev \ + file \ pkg-config \ git \ libidn11-dev \ @@ -46,7 +48,9 @@ if [ "${YDB_DEB_INSTALL_DEPS:-1}" = "1" ]; then python3 \ python3-six \ ragel \ - yasm + yasm \ + odbcinst \ + unixodbc-dev fi touch_existing_sources() { @@ -107,6 +111,7 @@ touch_existing_sources \ include \ library \ plugins \ + odbc \ scripts/build_cpack_deb_packages.sh \ scripts/generate-debian-directory.sh \ src \ @@ -121,8 +126,11 @@ cmake -S . -B build-deb \ -DYDB_SDK_TESTS=OFF \ -DYDB_SDK_ENABLE_OTEL_METRICS=ON \ -DYDB_SDK_ENABLE_OTEL_TRACE=ON \ + -DYDB_SDK_ODBC=ON \ -DBUILD_SHARED_LIBS=OFF \ -DYDB_SDK_USE_SYSTEM_GOOGLEAPIS=ON \ + -DYDB_ODBC_INSTALL_LIBDIR="/usr/lib/$(dpkg-architecture -qDEB_HOST_MULTIARCH)" \ + -DYDB_ODBC_INSTALL_DATADIR=/usr/share/ydb-odbc \ -DCMAKE_INSTALL_PREFIX=/usr/share/yandex \ -DCMAKE_PREFIX_PATH="/usr/share/yandex" \ "${CMAKE_COMPILER_LAUNCHER_ARGS[@]}" diff --git a/scripts/googleapis_deb/CMakeLists.txt b/scripts/googleapis_deb/CMakeLists.txt index 0c96c2358d4..b0cf2700666 100644 --- a/scripts/googleapis_deb/CMakeLists.txt +++ b/scripts/googleapis_deb/CMakeLists.txt @@ -48,6 +48,8 @@ endforeach() add_library(api-common-protos STATIC ${PROTO_SRCS} ${PROTO_HDRS}) add_library(yandex-googleapis-api-common-protos::api-common-protos ALIAS api-common-protos) +set_target_properties(api-common-protos PROPERTIES POSITION_INDEPENDENT_CODE ON) + target_include_directories(api-common-protos PUBLIC $ $ diff --git a/scripts/test_deb_packages.sh b/scripts/test_deb_packages.sh index 194c6090d8f..2c1e4341949 100755 --- a/scripts/test_deb_packages.sh +++ b/scripts/test_deb_packages.sh @@ -1,5 +1,5 @@ #!/bin/bash -set -e +set -euo pipefail if [ "$#" -ne 1 ]; then echo "Usage: $0 " @@ -10,17 +10,47 @@ DEB_DIR=$(realpath "$1") SCRIPT_DIR=$(dirname "$(realpath "$0")") SOURCE_DIR=$(realpath "$SCRIPT_DIR/..") TEST_DIR=$(realpath "$SCRIPT_DIR/../tests/deb_package") +YDB_TEST_IMAGE="${YDB_TEST_IMAGE:-ydbplatform/local-ydb:25.2.1}" +YDB_TEST_CONTAINER="ydb-odbc-package-test-$$" + +cleanup() { + docker rm -f "$YDB_TEST_CONTAINER" >/dev/null 2>&1 || true +} +trap cleanup EXIT echo "Building test Docker image..." docker build -t ydb-cpp-sdk-deb-test "$TEST_DIR" +echo "Starting local YDB ${YDB_TEST_IMAGE}..." +docker run -d --name "$YDB_TEST_CONTAINER" --network host \ + -e GRPC_TLS_PORT=2135 \ + -e GRPC_PORT=2136 \ + -e MON_PORT=8765 \ + -e YDB_DEFAULT_LOG_LEVEL=NOTICE \ + -e YDB_USE_IN_MEMORY_PDISKS=true \ + "$YDB_TEST_IMAGE" >/dev/null + +for _ in $(seq 1 60); do + if docker exec "$YDB_TEST_CONTAINER" /bin/sh -c \ + "/ydb -e grpc://localhost:2136 -d /local scheme ls" >/dev/null 2>&1; then + break + fi + sleep 2 +done +if ! docker exec "$YDB_TEST_CONTAINER" /bin/sh -c \ + "/ydb -e grpc://localhost:2136 -d /local scheme ls" >/dev/null 2>&1; then + docker logs "$YDB_TEST_CONTAINER" || true + echo "Local YDB did not become ready" >&2 + exit 1 +fi + echo "Running test container..." -docker run --rm \ +docker run --rm --network host \ -v "$DEB_DIR:/deb_packages:ro" \ -v "$SOURCE_DIR:/source:ro" \ ydb-cpp-sdk-deb-test \ bash -c ' -set -e +set -euo pipefail apt-get update if ! compgen -G "/deb_packages/yandex-googleapis-api-common-protos*.deb" > /dev/null; then @@ -34,6 +64,116 @@ else dpkg -i /deb_packages/yandex-googleapis-api-common-protos*.deb fi +odbc_packages=(/deb_packages/ydb-odbc_*.deb) +if [ "${#odbc_packages[@]}" -ne 1 ] || [ ! -f "${odbc_packages[0]}" ]; then + echo "Expected exactly one ydb-odbc package, found: ${odbc_packages[*]}" >&2 + exit 1 +fi +odbc_deb="${odbc_packages[0]}" +sdk_version="$(sed -nE '\''s/.*YDB_SDK_VERSION = "([0-9]+\.[0-9]+\.[0-9]+)".*/\1/p'\'' /source/src/version.h)" +package_version="$(dpkg-deb -f "$odbc_deb" Version)" +package_name="$(dpkg-deb -f "$odbc_deb" Package)" +package_arch="$(dpkg-deb -f "$odbc_deb" Architecture)" +package_depends="$(dpkg-deb -f "$odbc_deb" Depends)" +host_arch="$(dpkg --print-architecture)" +multiarch="$(dpkg-architecture -qDEB_HOST_MULTIARCH)" +driver_path="/usr/lib/${multiarch}/libydb-odbc.so" +driver_template="/usr/share/ydb-odbc/odbcinst.ini" + +test "$package_name" = ydb-odbc +test "$package_version" = "$sdk_version" +test "$package_arch" = "$host_arch" +for dependency in odbcinst libodbcinst2 libc6; do + if ! grep -Eq "(^|, )${dependency}([ (]|,|$)" <<<"$package_depends"; then + echo "Missing ydb-odbc dependency ${dependency}: ${package_depends}" >&2 + exit 1 + fi +done + +cat >/tmp/unrelated-odbcinst.ini </etc/odbc.ini </root/.odbc.ini </tmp/odbc-ini.sha256 + +verify_ydb_registration() { + local registration + registration="$(odbcinst -q -d -n YDB)" + grep -Fx "Driver=${driver_path}" <<<"$registration" + grep -Fx "Setup=${driver_path}" <<<"$registration" + grep -Fx "UsageCount=1" <<<"$registration" +} + +run_odbc_consumers() { + local isql_output + isql_output="$(printf "SELECT 42 AS value;\n" | isql -b -v YDBPackageTest)" + echo "$isql_output" + grep -Eq "(^|[^0-9])42([^0-9]|$)" <<<"$isql_output" + /odbc_qt_test/build/ydb_odbc_qt_test \ + "Driver={YDB};Server=localhost:2136;Database=/local" +} + +rm -rf /tmp/ydb-odbc-old +dpkg-deb --raw-extract "$odbc_deb" /tmp/ydb-odbc-old +sed -i "s/^Version: .*/Version: ${sdk_version}~package-test1/" \ + /tmp/ydb-odbc-old/DEBIAN/control +dpkg-deb --build /tmp/ydb-odbc-old /tmp/ydb-odbc-old.deb + +apt-get install -y /tmp/ydb-odbc-old.deb +test -f "$driver_path" +test -f "$driver_template" +grep -Fx "Driver=${driver_path}" "$driver_template" +verify_ydb_registration +run_odbc_consumers + +apt-get install -y "$odbc_deb" +test "$(dpkg-query -W -f='\''${Version}'\'' ydb-odbc)" = "$sdk_version" +verify_ydb_registration +run_odbc_consumers +sha256sum --check /tmp/odbc-ini.sha256 + +apt-get remove -y ydb-odbc +test ! -e "$driver_path" +test ! -e "$driver_template" +if odbcinst -q -d -n YDB >/dev/null 2>&1; then + echo "YDB remained registered after package removal" >&2 + exit 1 +fi +odbcinst -q -d -n UnrelatedPackageTest >/dev/null +sha256sum --check /tmp/odbc-ini.sha256 + +apt-get install -y "$odbc_deb" +verify_ydb_registration +odbcinst -u -d -n YDB +cat >/tmp/replacement-ydb-odbcinst.ini < +#include +#include +#include +#include + +#include + +int main(int argc, char** argv) { + QCoreApplication app(argc, argv); + + if (argc != 2) { + std::cerr << "Usage: " << argv[0] << " " << std::endl; + return 2; + } + + if (!QSqlDatabase::isDriverAvailable("QODBC")) { + std::cerr << "Qt QODBC plugin is not available" << std::endl; + return 3; + } + + const QString connectionName = QStringLiteral("ydb-odbc-package-test"); + { + QSqlDatabase database = QSqlDatabase::addDatabase("QODBC", connectionName); + database.setDatabaseName(QString::fromLocal8Bit(argv[1])); + + if (!database.open()) { + std::cerr << "QODBC connection failed: " + << database.lastError().text().toStdString() << std::endl; + return 4; + } + + QSqlQuery query(database); + query.setForwardOnly(true); + if (!query.exec(QStringLiteral("SELECT 42 AS value"))) { + std::cerr << "QODBC query failed: " + << query.lastError().text().toStdString() << std::endl; + return 5; + } + if (!query.next()) { + std::cerr << "QODBC returned no row: " + << query.lastError().text().toStdString() << std::endl; + return 6; + } + const QVariant value = query.value(0); + if (value.toInt() != 42) { + std::cerr << "QODBC returned an unexpected SELECT result: type=" + << value.typeName() << ", value=" + << value.toString().toStdString() << std::endl; + return 7; + } + + database.close(); + } + QSqlDatabase::removeDatabase(connectionName); + + return 0; +} diff --git a/tests/unit/library/operation_id/CMakeLists.txt b/tests/unit/library/operation_id/CMakeLists.txt index a6f2143949a..63d77da600a 100644 --- a/tests/unit/library/operation_id/CMakeLists.txt +++ b/tests/unit/library/operation_id/CMakeLists.txt @@ -5,6 +5,7 @@ add_ydb_test(NAME operation_id_ut GTEST yutil lib-operation_id-protos library-operation_id + cpp-testing-unittest LABELS unit )