diff --git a/.ci/assets/nixl-version-info.template b/.ci/assets/nixl-version-info.template new file mode 100644 index 0000000000..0de194b359 --- /dev/null +++ b/.ci/assets/nixl-version-info.template @@ -0,0 +1,21 @@ +# NIXL Container Version Information + +# Version information +NIXL_VERSION="${NIXL_VERSION}" +UCX_VERSION="${UCX_VERSION}" +EFA_INSTALLER_VERSION="${EFA_INSTALLER_VERSION}" + +# Build information +BUILD_TARGET="${BUILD_TARGET}" +BUILD_TIMESTAMP="${BUILD_TIMESTAMP}" +ARCHITECTURE="${arch}" +BASE_IMAGE_NAME="${BASE_IMAGE}" +BASE_IMAGE_TAG="${BASE_IMAGE_TAG}" +TAG_NAME="${TAG_NAME}" + +# Jenkins build information +BUILD_NUMBER="${BUILD_NUMBER}" +BUILD_URL="${BUILD_URL}" +JOB_NAME="${JOB_NAME}" +NODE_NAME="${NODE_NAME}" +WORKSPACE="${WORKSPACE}" diff --git a/.ci/assets/nixlbench-version-info.json.template b/.ci/assets/nixlbench-version-info.json.template deleted file mode 100644 index 586b188b71..0000000000 --- a/.ci/assets/nixlbench-version-info.json.template +++ /dev/null @@ -1,19 +0,0 @@ -{ - "versions": { - "nixl_version": "${NIXL_VERSION}", - "ucx_version": "${UCX_VERSION}" - }, - "build_info": { - "timestamp": "${BUILD_TIMESTAMP}", - "architecture": "${arch}", - "base_image_name": "${BASE_IMAGE}", - "base_image_tag": "${BASE_IMAGE_TAG}" - }, - "jenkins": { - "build_number": "${BUILD_NUMBER}", - "build_url": "${BUILD_URL}", - "job_name": "${JOB_NAME}", - "node_name": "${NODE_NAME}", - "workspace": "${WORKSPACE}" - } -} diff --git a/.ci/docs/setup_nvidia_gpu_with_rdma_support_on_ubuntu.md b/.ci/docs/setup_nvidia_gpu_with_rdma_support_on_ubuntu.md index 1418ffd85c..d3c04122dc 100644 --- a/.ci/docs/setup_nvidia_gpu_with_rdma_support_on_ubuntu.md +++ b/.ci/docs/setup_nvidia_gpu_with_rdma_support_on_ubuntu.md @@ -31,6 +31,11 @@ sudo apt install nvidia-kernel-open- # e.g., nvidia-kernel-open-575 sudo reboot ``` +If running on nvlink hosts like DGX we should also install fabric manager +```bash +sudo apt install nvidia-fabricmanager- # should be same as kernel version nvidia-fabricmanager-575 +``` + Verify with `nvidia-smi`. Driver compatibility is critical for RDMA support[^1_1][^1_3]. diff --git a/.ci/jenkins/lib/build-container-matrix.yaml b/.ci/jenkins/lib/build-container-matrix.yaml new file mode 100644 index 0000000000..9425534156 --- /dev/null +++ b/.ci/jenkins/lib/build-container-matrix.yaml @@ -0,0 +1,210 @@ +# NIXL Container Build Configuration +# Builds and pushes NIXL and NIXLBench containers with configurable NIXL/UCX versions + +--- +job: nixl-ci-build-container + +# Build settings +failFast: false +timeout_minutes: 240 + +# Infrastructure +kubernetes: + cloud: il-ipp-blossom-prod + namespace: swx-media + limits: "{memory: 16Gi, cpu: 8000m}" + requests: "{memory: 8Gi, cpu: 4000m}" + +runs_on_dockers: + - { name: "podman-v5.0.2", url: "quay.io/podman/stable:v5.0.2", privileged: true } + +# Build matrix +matrix: + axes: + arch: + - x86_64 + - aarch64 + +# Configuration +env: + REGISTRY_HOSTESS: "urm.nvidia.com" + REGISTRY_REPO: "sw-nbu-swx-nixl-docker-local/verification" + LOCAL_TAG_BASE: "nixl-ci:build-" + MAIL_FROM: "jenkins@nvidia.com" + NPROC: "16" + +taskName: "${BUILD_TARGET}/${arch}/${axis_index}" + +credentials: + - credentialsId: 'svc-nixl-artifactory-token' + usernameVariable: 'ARTIFACTORY_USERNAME' + passwordVariable: 'ARTIFACTORY_PASSWORD' + +pipeline_start: + shell: action + module: groovy + run: | + def suffix = params.TAG_SUFFIX ? "-${params.TAG_SUFFIX}" : "" + def buildName = params.BUILD_TARGET + currentBuild.displayName += "-${buildName}-${params.NIXL_VERSION}-${params.UCX_VERSION}${suffix}" + env.ENABLE_NIXL_BUILD = params.BUILD_TARGET == 'nixl' ? 'true' : 'false' + env.ENABLE_NIXLBENCH_BUILD = params.BUILD_TARGET == 'nixlbench' ? 'true' : 'false' + echo "ENABLE_NIXL_BUILD: ${env.ENABLE_NIXL_BUILD}" + echo "ENABLE_NIXLBENCH_BUILD: ${env.ENABLE_NIXLBENCH_BUILD}" + echo "BUILD_TARGET: ${params.BUILD_TARGET}" + +# Build pipeline +steps: + - name: Prepare + run: | + # Setup podman and dependencies + rm -f /etc/containers/storage.conf + podman system reset -f || true + ln -sfT $(type -p podman) /usr/bin/docker + yum install -y git gettext + + - name: Build NIXLBench + enable: ${ENABLE_NIXLBENCH_BUILD} + run: | + # Clone UCX source for nixlbench + git clone https://github.com/openucx/ucx.git ucx-src + git -C ucx-src checkout "${UCX_VERSION}" + + "benchmark/nixlbench/contrib/build.sh" \ + --base-image "${BASE_IMAGE}" \ + --base-image-tag "${BASE_IMAGE_TAG}" \ + --tag "${LOCAL_TAG_BASE}${arch}" \ + --arch "${arch}" \ + --no-cache \ + --nixl "$WORKSPACE" \ + --ucx "$WORKSPACE/ucx-src" + + - name: Build NIXL + enable: ${ENABLE_NIXL_BUILD} + run: | + export UCX_REF="${UCX_VERSION}" + + "contrib/build-container.sh" \ + --base-image "${BASE_IMAGE}" \ + --base-image-tag "${BASE_IMAGE_TAG}" \ + --tag "${LOCAL_TAG_BASE}${arch}" \ + --arch "${arch}" \ + --no-cache + + - name: Add Version Info + run: | + # Extract standardized 8-char commit hash for UCX version info: + if [[ "$BUILD_TARGET" == "nixlbench" ]]; then + CLEAN_UCX=$(cd "$WORKSPACE/ucx-src" && git rev-parse --short=8 HEAD) + else + UCX_REF="$UCX_VERSION" + + # Hash? if yes truncate, else ls-remote then truncate + if [[ "$UCX_REF" =~ ^[a-f0-9]{8,40}$ ]]; then + CLEAN_UCX="${UCX_REF:0:8}" + else + CLEAN_UCX=$(git ls-remote https://github.com/openucx/ucx.git "$UCX_REF" | head -n1 | cut -c1-8) + fi + + # Verify + [[ -n "$CLEAN_UCX" ]] || { echo "ERROR: failed to resolve UCX_REF=$UCX_REF"; exit 1; } + fi + + # Calculate tag name + NIXL_VERSION="$(git rev-parse --short=8 HEAD)" + TAG_NAME="${BASE_IMAGE_TAG}-nixl-${NIXL_VERSION}-ucx-${CLEAN_UCX}-${arch}${TAG_SUFFIX:+-${TAG_SUFFIX}}" + + # Generate version info file from template + export BUILD_TIMESTAMP="$(date -u '+%Y-%m-%dT%H:%M:%SZ')" \ + BUILD_TARGET BASE_IMAGE BASE_IMAGE_TAG arch \ + BUILD_NUMBER BUILD_URL JOB_NAME NODE_NAME WORKSPACE \ + NIXL_VERSION UCX_VERSION="${CLEAN_UCX}" TAG_NAME + envsubst < .ci/assets/nixl-version-info.template > version-info + + # Add version info to the image + docker run -itd --name tempcontainer "${LOCAL_TAG_BASE}${arch}" + docker cp version-info "tempcontainer:/opt/nixl-version" + docker commit tempcontainer "${LOCAL_TAG_BASE}${arch}" + docker rm -f tempcontainer || true + + - name: Push + credentialsId: 'svc-nixl-artifactory-token' + run: | + source version-info + ARTIFACTORY_REGISTRY="${REGISTRY_HOSTESS}/${REGISTRY_REPO}/${BUILD_TARGET}" + ARTIFACTORY_API="https://${REGISTRY_HOSTESS}/artifactory/api/storage/${REGISTRY_REPO}/${BUILD_TARGET}" + + # Prepare image properties + IMAGE_PROPERTIES="BUILD_TARGET=${BUILD_TARGET};NIXL_VERSION=${NIXL_VERSION};UCX_VERSION=${UCX_VERSION};arch=${arch};" + IMAGE_PROPERTIES+="BUILD_NUMBER=${BUILD_NUMBER};JOB_NAME=${JOB_NAME};BUILD_URL=${BUILD_URL};NODE_NAME=${NODE_NAME};" + IMAGE_PROPERTIES+="BASE_IMAGE=${BASE_IMAGE};BASE_IMAGE_TAG=${BASE_IMAGE_TAG}" + + # Login to Artifactory + echo "$ARTIFACTORY_PASSWORD" | docker login "${REGISTRY_HOSTESS}" -u "$ARTIFACTORY_USERNAME" --password-stdin + + # Function to tag, push, and set properties + tag_push_set_properties() { + local target_tag="$1" + echo "Creating tag: ${target_tag}" + docker tag "${LOCAL_TAG_BASE}${arch}" "${ARTIFACTORY_REGISTRY}:${target_tag}" + docker push "${ARTIFACTORY_REGISTRY}:${target_tag}" + curl -H "Authorization: Bearer ${ARTIFACTORY_PASSWORD}" -X PUT \ + "${ARTIFACTORY_API}/${target_tag}?properties=${IMAGE_PROPERTIES}" + } + + # Always create standard tag + tag_push_set_properties "${TAG_NAME}" + + # Check if latest tag should be updated + if [[ "${UPDATE_LATEST}" == "true" ]]; then + tag_push_set_properties "${BASE_IMAGE_TAG}-${arch}-latest" + fi + + - name: Show Results + run: | + source version-info + echo "Image type built: ${BUILD_TARGET} (${arch})" + echo "Image pushed to: ${REGISTRY_HOSTESS}/${REGISTRY_REPO}/${BUILD_TARGET}:${TAG_NAME}" + if [[ "${UPDATE_LATEST}" == "true" ]]; then + echo "Latest tag updated: ${REGISTRY_HOSTESS}/${REGISTRY_REPO}/${BUILD_TARGET}:${BASE_IMAGE_TAG}-${arch}-latest" + fi + + echo -e "\nBuild config for manual repro:" + if [[ "${BUILD_TARGET}" == "nixlbench" ]]; then + echo "git clone https://github.com/openucx/ucx.git ucx-src && (cd ucx-src && git checkout ${UCX_VERSION})" + echo "benchmark/nixlbench/contrib/build.sh --base-image ${BASE_IMAGE} --base-image-tag ${BASE_IMAGE_TAG} --tag local-test-tag --arch ${arch} --no-cache --nixl \$WORKSPACE --ucx \$WORKSPACE/ucx-src" + else + echo "export UCX_REF=${UCX_VERSION}" + echo "contrib/build-container.sh --base-image ${BASE_IMAGE} --base-image-tag ${BASE_IMAGE_TAG} --tag local-test-tag --arch ${arch} --no-cache" + fi + +pipeline_stop: + shell: action + module: groovy + run: | + if (params.MAIL_TO) { + def jobStatus = currentBuild.result ?: 'SUCCESS' + def statusColor = jobStatus == 'SUCCESS' ? 'green' : 'red' + + def userName = currentBuild.rawBuild.getCause(hudson.model.Cause.UserIdCause)?.userName ?: 'schedule' + + mail( + from: env.MAIL_FROM, + to: params.MAIL_TO, + subject: "Job '${env.JOB_NAME} [${env.BUILD_NUMBER}]' - ${jobStatus}", + mimeType: 'text/html', + body: """ +

Started by: ${userName}

+

Status: ${jobStatus}

+

Job: ${env.JOB_NAME}

+

Build: #${env.BUILD_NUMBER}

+

Console Output: Full Log

+

Build Target: ${params.BUILD_TARGET}

+

NIXL Version: ${params.NIXL_VERSION}

+

UCX Version: ${params.UCX_VERSION}

+

Base Image: ${params.BASE_IMAGE}:${params.BASE_IMAGE_TAG}

+

Architectures: x86_64, aarch64

+ ${params.UPDATE_LATEST ? '

Latest tag updated: Yes

' : ''} + """ + ) + } diff --git a/.ci/jenkins/lib/build-matrix.yaml b/.ci/jenkins/lib/build-matrix.yaml index e615a798ed..304ac34cee 100644 --- a/.ci/jenkins/lib/build-matrix.yaml +++ b/.ci/jenkins/lib/build-matrix.yaml @@ -46,6 +46,8 @@ matrix: env: NIXL_INSTALL_DIR: /opt/nixl + TEST_TIMEOUT: 30 + NPROC: "16" steps: - name: Build @@ -59,14 +61,28 @@ steps: - name: Test CPP parallel: false + timeout: "${TEST_TIMEOUT}" run: | .gitlab/test_cpp.sh ${NIXL_INSTALL_DIR} - name: Test Python parallel: false + timeout: "${TEST_TIMEOUT}" run: | .gitlab/test_python.sh ${NIXL_INSTALL_DIR} + - name: Test Nixlbench + parallel: false + timeout: "${TEST_TIMEOUT}" + run: | + .gitlab/test_nixlbench.sh ${NIXL_INSTALL_DIR} + + - name: Test Rust + parallel: false + timeout: "${TEST_TIMEOUT}" + run: | + .gitlab/test_rust.sh ${NIXL_INSTALL_DIR} + - name: Build Docker Image parallel: false containerSelector: "{ name: 'podman.*' }" diff --git a/.ci/jenkins/lib/nixlbench-container-build-matrix.yaml b/.ci/jenkins/lib/nixlbench-container-build-matrix.yaml deleted file mode 100644 index 1375e8868a..0000000000 --- a/.ci/jenkins/lib/nixlbench-container-build-matrix.yaml +++ /dev/null @@ -1,133 +0,0 @@ -# NIXLBench Container Build Configuration -# Builds and pushes NIXLBench containers with configurable NIXL/UCX versions - ---- -job: nixl-ci-nixlbench-container-build - -# Build settings -failFast: false -timeout_minutes: 240 - -# Infrastructure -kubernetes: - cloud: il-ipp-blossom-prod - namespace: swx-media - limits: "{memory: 16Gi, cpu: 8000m}" - requests: "{memory: 8Gi, cpu: 4000m}" - -runs_on_dockers: - - name: "podman-v5.0.2" - url: "quay.io/podman/stable:v5.0.2" - category: 'tool' - privileged: true - -# Build matrix -matrix: - axes: - arch: - - x86_64 - - aarch64 - -# Configuration -env: - ARTIFACTORY_HOST: "urm.nvidia.com" - ARTIFACTORY_REPO_PATH: "sw-nbu-swx-nixl-docker-local/verification/nixlbench" - BUILD_SCRIPT: "benchmark/nixlbench/contrib/build.sh" - LOCAL_TAG_BASE: "nixlbench:build-" - -credentials: - - credentialsId: 'svc-nixl-artifactory-token' - usernameVariable: 'ARTIFACTORY_USERNAME' - passwordVariable: 'ARTIFACTORY_PASSWORD' - -pipeline_start: - shell: action - module: groovy - run: | - def suffix = params.TAG_SUFFIX ? "-${params.TAG_SUFFIX}" : "" - currentBuild.displayName += "-${params.NIXL_VERSION}-${params.UCX_VERSION}${suffix}" - -# Build pipeline -steps: - - name: Prepare - parallel: false - containerSelector: "{ name: 'podman.*' }" - run: | - # Setup podman and dependencies - rm -f /etc/containers/storage.conf - podman system reset -f || true - ln -sfT $(type -p podman) /usr/bin/docker - yum install -y git gettext - - # Clone UCX source - git clone https://github.com/openucx/ucx.git ucx-src - cd ucx-src && git checkout "${UCX_VERSION}" - - - name: Build - parallel: false - containerSelector: "{ name: 'podman.*' }" - run: | - LOCAL_TAG="${LOCAL_TAG_BASE}${arch}" - - "${BUILD_SCRIPT}" \ - --nixl "$PWD" \ - --ucx "$PWD/ucx-src" \ - --base-image "${BASE_IMAGE}" \ - --base-image-tag "${BASE_IMAGE_TAG}" \ - --tag "${LOCAL_TAG}" \ - --arch "${arch}" \ - --no-cache - - # Generate version info file from template - export BUILD_TIMESTAMP="$(date -u '+%Y-%m-%dT%H:%M:%SZ')" \ - NIXL_VERSION UCX_VERSION BASE_IMAGE BASE_IMAGE_TAG arch \ - BUILD_NUMBER BUILD_URL JOB_NAME NODE_NAME WORKSPACE - envsubst < .ci/assets/nixlbench-version-info.json.template > version-info.json - - # Add version info to the image - CONTAINER_ID=$(docker create "${LOCAL_TAG}") - docker cp version-info.json "${CONTAINER_ID}:/opt/nixlbench-version.json" - docker commit "${CONTAINER_ID}" "${LOCAL_TAG}" - - - name: Push - parallel: false - containerSelector: "{ name: 'podman.*' }" - credentialsId: 'svc-nixl-artifactory-token' - run: | - LOCAL_TAG="${LOCAL_TAG_BASE}${arch}" - ARTIFACTORY_REGISTRY="${ARTIFACTORY_HOST}/${ARTIFACTORY_REPO_PATH}" - - # Sanitize versions for Docker tag compatibility (replace / and truncate) - CLEAN_NIXL=${NIXL_VERSION//\//-} - CLEAN_NIXL=${CLEAN_NIXL:0:14} - CLEAN_UCX=${UCX_VERSION//\//-} - CLEAN_UCX=${CLEAN_UCX:0:8} - - # Login to Artifactory - echo "$ARTIFACTORY_PASSWORD" | docker login "${ARTIFACTORY_REGISTRY}" -u "$ARTIFACTORY_USERNAME" --password-stdin - - # Prepare Artifactory API and image metadata - ARTIFACTORY_API="https://${ARTIFACTORY_HOST}/artifactory/api/storage/${ARTIFACTORY_REPO_PATH}" - IMAGE_PROPERTIES="NIXL_VERSION=${NIXL_VERSION};UCX_VERSION=${UCX_VERSION};arch=${arch};" - IMAGE_PROPERTIES+="BUILD_NUMBER=${BUILD_NUMBER};JOB_NAME=${JOB_NAME};" - IMAGE_PROPERTIES+="BUILD_URL=${BUILD_URL};NODE_NAME=${NODE_NAME};" - IMAGE_PROPERTIES+="BASE_IMAGE=${BASE_IMAGE};BASE_IMAGE_TAG=${BASE_IMAGE_TAG}" - - # Function to tag, push, and set properties - tag_push_set_properties() { - local target_tag="$1" - echo "Creating tag: ${target_tag}" - docker tag "${LOCAL_TAG}" "${ARTIFACTORY_REGISTRY}:${target_tag}" - docker push "${ARTIFACTORY_REGISTRY}:${target_tag}" - curl -H "Authorization: Bearer ${ARTIFACTORY_PASSWORD}" -X PUT \ - "${ARTIFACTORY_API}/${target_tag}?properties=${IMAGE_PROPERTIES}" - } - - # Always create standard tag: base-nixl-version-ucx-version-arch[-suffix] - TAG_NAME="${BASE_IMAGE_TAG}-nixl-${CLEAN_NIXL}-ucx-${CLEAN_UCX}-${arch}${TAG_SUFFIX:+-${TAG_SUFFIX}}" - tag_push_set_properties "${TAG_NAME}" - - # Check if latest tag should be updated - if [[ "${UPDATE_LATEST}" == "true" ]]; then - tag_push_set_properties "${BASE_IMAGE_TAG}-${arch}-latest" - fi diff --git a/.ci/jenkins/lib/test-matrix.yaml b/.ci/jenkins/lib/test-matrix.yaml index 37226e3e29..c7e8bad977 100644 --- a/.ci/jenkins/lib/test-matrix.yaml +++ b/.ci/jenkins/lib/test-matrix.yaml @@ -25,6 +25,7 @@ timeout_minutes: 240 # label is defined at jenkins slave configuration, we want to run the job on a gpu agent and be able to esaly replace it without having to change this file runs_on_agents: - {nodeLabel: 'H100'} + # - {nodeLabel: 'DGX'} matrix: axes: @@ -33,9 +34,15 @@ matrix: arch: - x86_64 +taskName: "${name}/${arch}/${axis_index}" + env: INSTALL_DIR: ${WORKSPACE}/nixl_install - UCX_VERSION: v1.19.x + UCX_VERSION: v1.19.0 + EFA_INSTALLER_VERSION: latest + NPROC: "16" + # Manual timeout - ci-demo doesn't handle docker exec + TEST_TIMEOUT: 30 steps: - name: Get Environment Info @@ -87,7 +94,7 @@ steps: - name: Build parallel: false run: | - docker exec -w ${WORKSPACE} -e UCX_VERSION=${UCX_VERSION} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/build.sh ${INSTALL_DIR}" + docker exec -w ${WORKSPACE} -e UCX_VERSION=${UCX_VERSION} -e EFA_INSTALLER_VERSION=${EFA_INSTALLER_VERSION} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/build.sh ${INSTALL_DIR}" onfail: | docker rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" docker image rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" @@ -95,7 +102,7 @@ steps: - name: Test CPP parallel: false run: | - docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_cpp.sh ${INSTALL_DIR}" + timeout ${TEST_TIMEOUT}m docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_cpp.sh ${INSTALL_DIR}" onfail: | docker rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" docker image rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" @@ -103,7 +110,23 @@ steps: - name: Test Python parallel: false run: | - docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_python.sh ${INSTALL_DIR}" + timeout ${TEST_TIMEOUT}m docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_python.sh ${INSTALL_DIR}" + onfail: | + docker rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" + docker image rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" + + - name: Test Nixlbench + parallel: false + run: | + timeout ${TEST_TIMEOUT}m docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_nixlbench.sh ${INSTALL_DIR}" + onfail: | + docker rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" + docker image rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" + + - name: Test Rust + parallel: false + run: | + timeout ${TEST_TIMEOUT}m docker exec -w ${WORKSPACE} "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" /bin/bash -c ".gitlab/test_rust.sh ${INSTALL_DIR}" always: | docker rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" docker image rm -f "${JOB_BASE_NAME}-${BUILD_ID}-${axis_index}" diff --git a/.ci/jenkins/pipeline/proj-jjb.yaml b/.ci/jenkins/pipeline/proj-jjb.yaml index cf7640c185..f9c7735643 100644 --- a/.ci/jenkins/pipeline/proj-jjb.yaml +++ b/.ci/jenkins/pipeline/proj-jjb.yaml @@ -16,8 +16,8 @@ url: "{jjb_gh_url}" # Build history retention policy - build-discarder: - days-to-keep: 50 # Keep builds for 50 days - num-to-keep: 20 # Or keep last 20 builds, whichever comes first + days-to-keep: 14 # Keep builds for 14 days + num-to-keep: 1000 # Or keep last 1000 builds, whichever comes first # Inject project-specific variables - inject: keep-system-variables: true @@ -96,15 +96,15 @@ - job-template: name: "{jjb_proj}-build" # Will be expanded to 'nixl-ci-build' project-type: pipeline - disabled: false folder: "{jjb_folder}" + disabled: false properties: # Similar properties as dispatcher job - github: url: "{jjb_gh_url}" - build-discarder: - days-to-keep: 50 - num-to-keep: 20 + days-to-keep: 14 + num-to-keep: 1000 - inject: keep-system-variables: true properties-content: | @@ -166,8 +166,8 @@ - github: url: "{jjb_gh_url}" - build-discarder: - days-to-keep: 50 - num-to-keep: 20 + days-to-keep: 14 + num-to-keep: 1000 - inject: keep-system-variables: true properties-content: | @@ -218,14 +218,14 @@ parent-credentials: true script-path: "{jjb_jenkinsfile}" # Path to Jenkinsfile that defines the build steps -# Template for NIXLBench container build job -# Builds and pushes NIXLBench container images for x86_64 and aarch64 +# Template for NIXL build container job +# Builds and pushes NIXL and NIXLBench container images for x86_64 and aarch64 # Supports nightly automatic builds and manual builds via parameters - job-template: - name: "{jjb_proj}-nixlbench-container-build" + name: "{jjb_proj}-build-container" project-type: pipeline - disabled: false folder: "{jjb_folder}" + disabled: false properties: - build-discarder: days-to-keep: 30 @@ -233,29 +233,44 @@ - inject: keep-system-variables: true properties-content: | - jjb_proj={jjb_proj}-nixlbench-container-build - conf_file=.ci/jenkins/lib/nixlbench-container-build-matrix.yaml + jjb_proj={jjb_proj}-build-container + conf_file=.ci/jenkins/lib/build-container-matrix.yaml description: > - NIXLBench container build
+ NIXL container build
• Builds and pushes x86_64 & aarch64 images with any NIXL/UCX version combination
+ • Choose between nixlbench or nixl build targets
• Optional latest tag update via UPDATE_LATEST parameter
- • All images pushed to unified path: verification/nixlbench/
+ • Images pushed to: verification/nixlbench/ or verification/nixl/ based on target

Do NOT edit this job through the Jenkins GUI — managed by Jenkins Job Builder. - concurrent: false + concurrent: true sandbox: true # Nightly scheduler for automatic builds with default versions triggers: - - timed: "H 3 * * *" # Run nightly around 3 AM (H adds randomness within the hour) + - parameterized-timer: + cron: | + # Build nixlbench nightly around 3 AM - Ubuntu 24.04 + H 3 * * * %BUILD_TARGET=nixlbench;UPDATE_LATEST=true + # Build nixl nightly around 4 AM - Ubuntu 24.04 + H 4 * * * %BUILD_TARGET=nixl;UPDATE_LATEST=true # Manual build parameters parameters: + - choice: + name: "BUILD_TARGET" + choices: + - "nixlbench" + - "nixl" + description: > + Build target:
+ • nixlbench: Builds NIXLBench container with benchmark tools
+ • nixl: Builds NIXL library container - string: name: "NIXL_VERSION" default: "{jjb_branch}" - description: "NIXL version to use (tag like 0.5.0, branch name, or commit hash)" + description: "NIXL version to use (tag like 0.6.0, branch name, or commit hash)" - string: name: "UCX_VERSION" - default: "v1.19.x" + default: "v1.19.0" description: "UCX version to use (tag like v1.19.0, branch name, or commit hash)" - string: name: "BASE_IMAGE" @@ -269,8 +284,8 @@ name: "TAG_SUFFIX" default: "" description: > - Optional tag suffix. Does not apply if update latest is set. Tag format:
- <base-image-tag>-nixl-<nixl-version>-ucx-<ucx-version>-<arch>[-<suffix>]
+ Optional tag suffix. Tag format:
+ <base-image-tag>-<nixl-version>-ucx-<ucx-version>-<arch>[-<suffix>]
- bool: name: "UPDATE_LATEST" default: false @@ -278,6 +293,10 @@ Update the latest tag for this architecture.
When enabled, also creates: <base-image-tag>-<arch>-latest
Example: 25.03-cuda12.8-devel-ubuntu24.04-aarch64-latest
+ - string: + name: "MAIL_TO" + default: "25f58ae0.NVIDIA.onmicrosoft.com@amer.teams.ms" + description: "Email address to send build results (optional)" # SCM configuration pipeline-scm: scm: @@ -310,5 +329,5 @@ jobs: - "{jjb_proj}-dispatcher" # Create dispatcher job - "{jjb_proj}-build" # Create build job - - "{jjb_proj}-nixlbench-container-build" # Create NIXLBench container build job + - "{jjb_proj}-build-container" # Create container builder job - "{jjb_proj}-test" # Create test job diff --git a/.ci/scripts/check_prints.sh b/.ci/scripts/check_prints.sh new file mode 100755 index 0000000000..33d52a754e --- /dev/null +++ b/.ci/scripts/check_prints.sh @@ -0,0 +1,60 @@ +#!/bin/bash +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Check if a path is provided as an argument +if [ -z "$1" ]; then + echo "Usage: $0 " + echo "Example: $0 ./path/to/python/directory" + exit 1 +fi + +DIR_PATH="$1" + +# Validate that the provided path is a directory +if [ ! -d "$DIR_PATH" ]; then + echo "Error: The provided path '$DIR_PATH' is not a valid directory." + exit 1 +fi + +echo "Checking for BUILT-IN 'print()' calls in Python files within: $DIR_PATH" +echo "---------------------------------------------------------------------" + +found_print=false + +# Find all Python files and process them +while read -r py_file; do + # Use grep to find 'print()' calls with line numbers, then filter out method calls. + # First grep: finds all occurrences of 'print(' with word boundary. + # Second grep: filters out lines where 'print(' is preceded by a dot and optional whitespace. + MATCHES=$(grep -nE '\bprint\s*\(' "$py_file" | grep -vE '\.[[:space:]]*print\s*\(') + + if [ -n "$MATCHES" ]; then + echo "Found built-in 'print()' in: $py_file" + echo "${MATCHES//$'\n'/$'\n' Line }" # Indent and prepend "Line " + echo # Add a blank line for readability + found_print=true + fi +done < <(find "$DIR_PATH" -name "*.py") + +echo "---------------------------------------------------------------------" + +if [ "$found_print" = true ]; then + echo "One or more Python files in '$DIR_PATH' contain built-in 'print()' calls." + exit 1 +else + echo "No built-in 'print()' calls found in any Python files within '$DIR_PATH'." + exit 0 +fi diff --git a/.ci/scripts/common.sh b/.ci/scripts/common.sh new file mode 100755 index 0000000000..f321e54d11 --- /dev/null +++ b/.ci/scripts/common.sh @@ -0,0 +1,76 @@ +#!/bin/bash +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# +# Common functions for CI scripts +# + +# +# Set initial port number for client/server applications to be updated with +# function below +# +tcp_port_range=1000 +min_port_number=10500 +max_port_number=65535 + +# GITLAB CI +if [ -n "$CI_CONCURRENT_ID" ]; then + nixl_concurrent_id=$CI_CONCURRENT_ID +# Jenkins CI +elif [ -n "$EXECUTOR_NUMBER" ]; then + nixl_concurrent_id=$EXECUTOR_NUMBER +else + # Fallback to random number if both CI_CONCURRENT_ID and EXECUTOR_NUMBER are not set + nixl_concurrent_id=$((RANDOM % $(((max_port_number - min_port_number) / tcp_port_range)))) +fi + +echo nixl_concurrent_id="$nixl_concurrent_id" + +# First half of the port range is used for shell script tests +tcp_port_min=$((min_port_number + nixl_concurrent_id * tcp_port_range)) +tcp_port_max=$((tcp_port_min + tcp_port_range / 2)) + +get_next_tcp_port() { + local port_file="/tmp/nixl_tcp_port_${nixl_concurrent_id}" + + if [ ! -f "$port_file" ]; then + echo "$tcp_port_min" > "$port_file" + fi + + local current_port + current_port=$(cat "$port_file") + local next_port=$((current_port + 1)) + + # Check if the port is already in use + while ss -tuln | grep -q :$next_port; do + next_port=$((next_port + 1)) + done + + if [ "$next_port" -ge "$tcp_port_max" ]; then + next_port="$tcp_port_min" + fi + + echo "$next_port" > "$port_file" + + echo "$next_port" +} + +# Second half of the port range is used for gtest +gtest_offset=$((tcp_port_range / 2)) +# shellcheck disable=SC2034 +min_gtest_port=$((tcp_port_min + gtest_offset)) +# shellcheck disable=SC2034 +max_gtest_port=$((tcp_port_max + gtest_offset)) diff --git a/.github/workflows/aws_efa_validation.yml b/.github/workflows/aws_efa_validation.yml index a8448325ac..e8456ebbbc 100644 --- a/.github/workflows/aws_efa_validation.yml +++ b/.github/workflows/aws_efa_validation.yml @@ -15,6 +15,8 @@ jobs: AWS_DEFAULT_REGION: eu-central-1 AWS_ACCESS_KEY_ID: ${{ secrets.AWS_ACCESS_KEY_ID }} AWS_SECRET_ACCESS_KEY: ${{ secrets.AWS_SECRET_ACCESS_KEY }} + # tmp workaround: pin UCX ver to 1.18 until EFA support is available in UCX 1.19. + UCX_VERSION: v1.18.0 strategy: fail-fast: false matrix: @@ -25,6 +27,9 @@ jobs: - test_name: "Python Tests" test_scripts: - .gitlab/test_python.sh $NIXL_INSTALL_DIR + - test_name: "Rust Tests" + test_scripts: + - .gitlab/test_rust.sh $NIXL_INSTALL_DIR steps: - name: Checkout repository @@ -53,6 +58,8 @@ jobs: - name: Run AWS tests working-directory: ./contrib/aws-efa timeout-minutes: 180 + env: + TEST_TIMEOUT: 30 run: | set -o pipefail test_cmd='${{ join(matrix.test_scripts, ' && ') }}' diff --git a/.github/workflows/blossom-ci.yml b/.github/workflows/blossom-ci.yml index 42586944a9..6c21f39004 100644 --- a/.github/workflows/blossom-ci.yml +++ b/.github/workflows/blossom-ci.yml @@ -3,6 +3,11 @@ name: Blossom-CI on: issue_comment: types: [created] + pull_request: + types: [opened, reopened, synchronize] + branches: + - 'release/*' + - main workflow_dispatch: inputs: platform: diff --git a/.github/workflows/copyright-check.ps1 b/.github/workflows/copyright-check.ps1 deleted file mode 100644 index 77efd573f0..0000000000 --- a/.github/workflows/copyright-check.ps1 +++ /dev/null @@ -1,426 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2022-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -set-strictmode -version latest - -. "$(& git rev-parse --show-toplevel)/.github/workflows/common.ps1" - -# == begin common.ps1 extensions == - -$date_key = '%%DATE%%' -$date_regex = '(?>(?>\d{4})-)?(?\d{4})' - -$timer = [System.Diagnostics.Stopwatch]::StartNew() - -$global:copyright_matchers = @( - @{ - files = @('.containerfile', '.dockerignore', '.pbtxt', '.ps1', '.py', '.sh', '.toml', '.tpl', '.txt', '.yaml', '.yml', 'Dockerfile', '.build') - found_missing = $false - matches = @( - '# SPDX-FileCopyrightText: Copyright (c) ' + $date_key + ' NVIDIA CORPORATION & AFFILIATES. All rights reserved.' - '# SPDX-License-Identifier: Apache-2.0' - '# Licensed under the Apache License, Version 2.0 (the "License");' - '# you may not use this file except in compliance with the License.' - '# You may obtain a copy of the License at' - '# http://www.apache.org/licenses/LICENSE-2.0' - '# Unless required by applicable law or agreed to in writing, software' - '# distributed under the License is distributed on an "AS IS" BASIS,' - '# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.' - '# See the License for the specific language governing permissions and' - '# limitations under the License.' - ) - name = 'basic' - regex = $null - vertical_spacer = '#' - } - @{ - files = @('.cpp', '.h') - found_missing = $false - matches = @( - '/*' - ' * SPDX-FileCopyrightText: Copyright (c) ' + $date_key + ' NVIDIA CORPORATION & AFFILIATES. All rights reserved.' - ' * SPDX-License-Identifier: Apache-2.0' - ' *' - ' * Licensed under the Apache License, Version 2.0 (the "License");' - ' * you may not use this file except in compliance with the License.' - ' * You may obtain a copy of the License at' - ' *' - ' * http://www.apache.org/licenses/LICENSE-2.0' - ' *' - ' * Unless required by applicable law or agreed to in writing, software' - ' * distributed under the License is distributed on an "AS IS" BASIS,' - ' * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.' - ' * See the License for the specific language governing permissions and' - ' * limitations under the License.' - ' */' - ) - name = 'c-code' - regex = $null - vertical_spacer = ' *' - } - @{ - files = @('.json') - found_missing = $false - matches = @( - '"copyright": [' - ' "SPDX-FileCopyrightText: Copyright (c) ' + $date_key + ' NVIDIA CORPORATION & AFFILIATES. All rights reserved.",' - ' "SPDX-License-Identifier: Apache-2.0",' - ' "Licensed under the Apache License, Version 2.0 (the \"License\");",' - ' "you may not use this file except in compliance with the License.",' - ' "You may obtain a copy of the License at",' - ' "http://www.apache.org/licenses/LICENSE-2.0",' - ' "Unless required by applicable law or agreed to in writing, software",' - ' "distributed under the License is distributed on an \"AS IS\" BASIS,",' - ' "WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.",' - ' "See the License for the specific language governing permissions and",' - ' "limitations under the License."' - '],' - ) - name = 'json' - regex = $null - vertical_spacer = $null - } - @{ - files = @('.md') - found_missing = $false - matches = @( - '' - ) - name = 'markdown' - regex = $null - vertical_spacer = '' - } - @{ - files = @('.proto', '.rs') - found_missing = $false - matches = @( - '// SPDX-FileCopyrightText: Copyright (c) ' + $date_key + ' NVIDIA CORPORATION & AFFILIATES. All rights reserved.' - '// SPDX-License-Identifier: Apache-2.0' - '// Licensed under the Apache License, Version 2.0 (the "License");' - '// you may not use this file except in compliance with the License.' - '// You may obtain a copy of the License at' - '// http://www.apache.org/licenses/LICENSE-2.0' - '// Unless required by applicable law or agreed to in writing, software' - '// distributed under the License is distributed on an "AS IS" BASIS,' - '// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.' - '// See the License for the specific language governing permissions and' - '// limitations under the License.' - ) - name = 'c-like' - regex = $null - vertical_spacer = '//' - } -) -$global:copyright_results = @{ - failed_date = @() - failed_header = @() - passed = @() - skipped = @() - unsupported = @() -} - -# === end common.ps1 extensions === - -$ignored_files = @('.clang-format', '.gitattributes', '.gitignore', '.gitkeep', '.patch', 'Cargo.lock', 'LICENSE', 'uv.lock', 'rust-toolchain.toml', 'codespell.txt', 'aws_job_def.json') -write-debug " ignored_files = ['$($ignored_files -join "','")']." -$ignored_paths = @('.github', '.mypy_cache', '.pytest_cache', '.ci/jenkins') -write-debug " ignored_paths = ['$($ignored_paths -join "','")']." -$ignored_types = @('.bat', '.gif', '.ico', '.ipynb', '.jpg', '.jpeg', '.patch', '.png', '.pyc', '.pyi', '.rst', '.zip', '.md') -write-debug " ignored_types = ['$($ignored_types -join "', '")']." -$ignored_folders = @('.git', '__pycache__') - -function is_ignored([string] $path) { - # write-debug " path: \"${path}\"." - if (($null -eq $path) -or ($path.length -eq 0) -or ($path.endswith('/'))) { - write-debug " ignored: true." - return $true - } - - foreach ($ignored_path in $ignored_paths) { - if ($path.startswith($ignored_path)) { - write-debug " ignored: true." - return $true - } - } - - foreach ($ignored_extension in $ignored_types) { - if ($file.endswith($ignored_extension)) { - write-debug " ignore = true." - return $true - } - } - - foreach ($ignored_file in $ignored_files) { - if ($file.endswith($ignored_file)) { - write-debug " ignore = true." - return $true - } - } - - $normalized_path = normalize_path $file - - foreach ($ignored_folder in $ignored_folders) { - if ($normalized_path -contains "/${ignored_folder}/") { - write-debug " ignore = true." - return $true - } - } - - if (-not(test-path "${normalized_path}" -pathtype 'Leaf')) { - write-debug " ignore = true." - return $true - } - - write-debug " ignore = false." - return $false -} - -function build_regex([object] $matcher) { - write-debug " matcher.name: $($matcher.name)." - - $regex = '' - foreach ($match in $matcher.matches) { - $match = $match -replace '([\(\)\[\]\.\+\*\\])', '\$1' - $match = $match -replace '\s+', '\s+' - # Given the amount of inconsistency between using http and https, we'll just regex it away. - $match = $match -replace 'https?://', 'https?://' - # Replace the date matcher placeholder w/ the actual regex we'll need. - $match = $match -replace $date_key, $date_regex - - $regex = "${regex}${match}[\n\r\s]+" - if ($null -ne $matcher.vertical_spacer) { - $regex = "${regex}(?>$($matcher.vertical_spacer)[\n\r\s]+)*" - } - } - - write-debug " -> '${regex}'." - return $regex -} - -function check_header([string] $path, [object] $matcher) { - write-debug " path: ""${path}""." - write-debug " matcher: ""$($matcher.name)""." - - $command = "git log -1 --pretty=""%cs"" -- ${file}" - $output = invoke-expression $command | out-string - $output = $output.trim() - $last_modified = $output.substring(0, 4) -as [int] - - write-debug " last_modified: ${last_modified}." - - if ($null -eq $matcher.regex) { - $matcher.regex = $(build_regex $matcher) - } - $regex = $matcher.regex - - write-debug " regex: ""${regex}""." - - $contents = read_content $path - if (($null -eq $contents) -or ($contents.length -le 0)) { - $global:copyright_results.skipped += $file - write-detailed " [SKIP] ${file}" 'DarkGray' - return - } - - if ($contents -match $regex) { - $capture_date = $Matches.year -as [int] - - if ($capture_date -lt $last_modified) { - $global:copyright_results.failed_date += $file - write-error " [FAIL] Incorrect Date in Header: ${path} (${capture_date})" - } - else { - $global:copyright_results.passed += $file - write-normal " [PASS] ${file}" - } - } - else { - $global:copyright_results.failed_header += $file - write-error " [FAIL] Invalid/Missing Header: ${file}" - $matcher.found_missing = $true - } -} - -function check_file([string] $file) { - write-debug " file: ""${file}""." - - $path = normalize_path $file - - if (test-path $path -pathtype 'Leaf') { - write-debug " path: ""${path}""." - - $is_checked = $false - foreach ($matcher in $global:copyright_matchers) - { - foreach ($ext in $matcher.files) { - if ($path.endswith($ext)) { - check_header $path $matcher - $is_checked = $true - break - } - } - if ($is_checked) { - return - } - } - - write-warning " [WARN] Unsupported: ${file}" - $global:copyright_results.unsupported += $file - } -} - -$current_year = "$(get-date -format 'yyyy')" -as [int] -write-debug " current_year = ${current_year}." - -foreach ($file in $(git ls-tree -r --name-only HEAD)) { - $file = $file.trim() - write-debug " file: ""${file}""." - - if (is_ignored $file) { - write-detailed " [SKIP] ${file}" 'DarkGray' - $global:copyright_results.skipped += $file - continue - } - - check_file $file -} - -function generate_report() { - $reports_path = $env:NVBUILD_REPORTS_PATH - if ($null -eq $reports_path) { - return - } - - if (-not (test-path $reports_path -pathtype 'Container')) { - if (test-path $reports_path -pathtype 'Leaf') { - return - } - new-item $reports_path -itemtype 'Directory' | out-null - } - - write-debug " Generating check report." - - $check_results = "`n" - - if ($global:copyright_results.failed_header.count -gt 0) { - $check_results += " `n" - - foreach ($file in $global:copyright_results.failed_header) { - $check_results += " ${file}`n" - } - - $check_results += " `n" - } - - if ($global:copyright_results.failed_date.count -gt 0) { - $check_results += " `n" - - foreach ($file in $global:copyright_results.failed_date) { - $check_results += " ${file}`n" - } - - $check_results += " `n" - } - - if ($global:copyright_results.unsupported.count -gt 0) { - $check_results += " `n" - - foreach ($file in $global:copyright_results.unsupported) { - $check_results += " ${file}`n" - } - - $check_results += " `n" - } - - if ($global:copyright_results.passed.count -gt 0) { - $check_results += " `n" - - foreach ($file in $global:copyright_results.passed) { - $check_results += " ${file}`n" - } - - $check_results += " `n" - } - - if ($global:copyright_results.skipped.count -gt 0) { - $check_results += " `n" - - foreach ($file in $global:copyright_results.skipped) { - $check_results += " ${file}`n" - } - - $check_results += " `n" - } - - $check_results += "`n" - $output_path = "${reports_path}/copyright-check.xml" - - write_content $check_results $output_path -overwrite - - write-minimal '' - write-minimal "Copyright check report -> ${output_path}" -} -write-normal '' - -$timer.Stop() - -write-high "Pass: $($global:copyright_results.passed.count), Fail: $($global:copyright_results.failed_date.count + $global:copyright_results.failed_header.count)" -no_newline -if ($global:copyright_results.skipped.count -gt 0) { - write-high ", Skipped: $($global:copyright_results.skipped.count)" -no_newline -} -if ($global:copyright_results.unsupported.count -gt 0) { - write-high ", Unsupported: $($global:copyright_results.unsupported.count)" -no_newline -} -write-minimal " ($($timer.Elapsed.TotalSeconds.ToString("0.000")) seconds)" $global:colors.low -no_newline -write-minimal '' - -if ($global:copyright_results.failed_header.count -gt 0) { - write-low '' - write-low 'Copyright checkers detected missing or invalid copyright headers:' - write-low '' - foreach ($matcher in $global:copyright_matchers) { - if ($matcher.found_missing) { - write-low " name: $($matcher.name)" - write-low " files: $($matcher.files -join ", ")" - write-low " pattern:`n $($matcher.regex)" - write-low '' - } - } -} - - -if (($global:copyright_results.failed_date.count -gt 0) -or ($global:copyright_results.failed_header.count -gt 0)) { - write-high 'Files out of compliance:' - # Final, end of output list of errors. - foreach ($path in $global:copyright_results.failed_header) { - write-error " [FAIL] invalid/missing header: ${path}" - } - foreach ($path in $global:copyright_results.failed_date) { - write-error " [FAIL] incorrect date: ${path}" - } - exit(-1) -} diff --git a/.github/workflows/copyright-check.sh b/.github/workflows/copyright-check.sh new file mode 100755 index 0000000000..1fe931983d --- /dev/null +++ b/.github/workflows/copyright-check.sh @@ -0,0 +1,72 @@ +#!/usr/bin/env bash +# SPDX-FileCopyrightText: Copyright (c) 2022-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +set -euo pipefail + +failures=() + +for f in $(git ls-files); do + # Normalize path + f=${f#./} + + # Skip ignored folders anywhere in path + case "$f" in + .github/*|.ci/*) + continue + ;; + esac + + # Skip ignored top-level paths + case "$f" in + *.png|*.jpg|*.jpeg|*.gif|*.ico|*.zip|*.rst|*.pyc|*.lock|*.md|*.svg|*.wrap|*.in|*.json|*.template|*.gitignore|*.python-version|*py.typed) + continue + ;; + CODEOWNERS|LICENSE|Doxyfile|.clang-format|.clang-tidy|.codespellrc) + continue + ;; + esac + + header=$(head -n 20 "$f") + + # Match SPDX-FileCopyrightText with NVIDIA and year(s) + if ! echo "$header" | grep -Eq 'SPDX-FileCopyrightText:\s*Copyright \(c\) [0-9]{4}(-[0-9]{4})? NVIDIA CORPORATION & AFFILIATES\. All rights reserved\.'; then + failures+=("$f (missing or incorrect copyright line)") + continue + fi + + # Extract last modification year from git + last_modified=$(git log -1 --pretty="%cs" -- "$f" | cut -d- -f1) + + # Extract copyright years (handles YYYY or YYYY-YYYY) + copyright_years=$(echo "$header" | \ + grep -Eo 'Copyright \(c\) [0-9]{4}(-[0-9]{4})?' | \ + sed -E 's/.* ([0-9]{4})(-[0-9]{4})?/\1\2/' || true) + + if [[ -z "$copyright_years" ]]; then + failures+=("$f (missing copyright)") + continue + fi + + # Get last year (handles range) + end_year=$(echo "$copyright_years" | sed -E 's/.*-//' || true) + + # Validate date + if (( end_year < last_modified )); then + failures+=("$f (copyright year $end_year < last modified $last_modified)") + continue + fi + + # License line must exist + if ! echo "$header" | grep -Eq '^[[:space:]]*(#|//|\*|/\*| # NVIDIA Inference Xfer Library (NIXL) @@ -46,14 +34,14 @@ pip install nixl ### UCX -NIXL was tested with UCX version 1.18.0. +NIXL was tested with UCX version 1.19.0. [GDRCopy](https://github.com/NVIDIA/gdrcopy) is available on Github and is necessary for maximum performance, but UCX and NIXL will work without it. ``` -$ wget https://github.com/openucx/ucx/releases/download/v1.18.0/ucx-1.18.0.tar.gz -$ tar xzf ucx-1.18.0.tar.gz -$ cd ucx-1.18.0 +$ wget https://github.com/openucx/ucx/releases/download/v1.19.0/ucx-1.19.0.tar.gz +$ tar xzf ucx-1.19.0.tar.gz +$ cd ucx-1.19.0 $ ./configure \ --enable-shared \ --disable-static \ @@ -170,15 +158,23 @@ pip install . For Python examples, see [examples/python/](examples/python/). ### Rust Bindings +#### Build +- Use `-Drust=true` meson option to build rust bindings. +- Use `-Ddebug=false` for a release build. +- Or build manually: + ```bash + $ cargo build --release + ``` +#### Install +The bindings will be installed under `nixl-sys` in the configured installation prefix. +Can be done using ninja, from project build directory: ```bash -# Build with default NIXL installation (/opt/nvidia/nvda_nixl) -$ cd src/bindings/rust -$ cargo build --release - -# Or specify custom NIXL location -$ NIXL_PREFIX=/path/to/nixl cargo build --release +$ ninja install +``` -# Run tests +#### Test +``` +# Rust bindings tests $ cargo test ``` diff --git a/benchmark/kvbench/README.md b/benchmark/kvbench/README.md index 3a71e52b16..d8214de562 100644 --- a/benchmark/kvbench/README.md +++ b/benchmark/kvbench/README.md @@ -118,7 +118,7 @@ These arguments are used by both `plan` and `profile` commands: | -------- | ----------- | | `--source` | Source of the nixl descriptors [file, memory, gpu] (default: file) | | `--destination` | Destination of the nixl descriptors [file, memory, gpu] (default: memory) | -| `--backend` | Communication backend [UCX, UCX_MO, GDS] (default: UCX) | +| `--backend` | Communication backend [UCX, UCX_MO, GDS, GDS_MT, POSIX, GPUNETIO, Mooncake, HF3FS, OBJ] (default: UCX) | | `--worker_type` | Worker to use to transfer data [nixl, nvshmem] (default: nixl) | | `--initiator_seg_type` | Memory segment type for initiator [DRAM, VRAM] (default: DRAM) | | `--target_seg_type` | Memory segment type for target [DRAM, VRAM] (default: DRAM) | @@ -137,6 +137,7 @@ These arguments are used by both `plan` and `profile` commands: | `--num_initiator_dev` | Number of devices in initiator processes (default: 1) | | `--num_target_dev` | Number of devices in target processes (default: 1) | | `--enable_pt` | Enable progress thread | +| `--progress_threads` | Number of progress threads (default: 0) | | `--device_list` | Comma-separated device names (default: all) | | `--runtime_type` | Type of runtime to use [ETCD] (default: ETCD) | | `--etcd-endpoints` | ETCD server URL for coordination (default: http://localhost:2379) | @@ -204,12 +205,14 @@ Benchmark the performance of a continuum of traffic patterns one after the other 0.129 2.147 2.386 4 ``` -#### CT Perftest +#### CT Perftest {#ct-perftest} Benchmark the performance of one traffic pattern. The pattern is run in multiple iterations and then metrics are reported. Useful for optimizing specific patterns. **Reports**: CT Perftest reports total latency (time elapsed between the first rank started until the last rank finished), average time per iteration, total size sent over the network, and average bandwidth by rank. +**Important note**: GPU memory is allocated with pytorch on the GPU specified by the `CUDA_VISIBLE_DEVICE` environment variable, make sure that each process sets this variable to the right device. + ## Examples ### KVBench Examples @@ -324,7 +327,7 @@ traffic_patterns: **Traffic Pattern Parameters**: - `matrix_file`: File containing the transfer matrix (required) - `shards`: Number of chunks to shard the buffer into (default: 1) -- `mem_type`: Memory type, currently supports "cuda" (default: "cuda") +- `mem_type`: Memory type, currently supports "cuda" (default: "cuda") (Use `CUDA_VISIBLE_DEVICES` to control the GPU device) - `xfer_op`: Transfer operation, "READ" or "WRITE" (default: "WRITE") - `sleep_after_launch_sec`: Seconds to sleep before running pattern (default: 0) @@ -359,6 +362,8 @@ python test/inference_workload_matgen.py generate \ #### Running CTP Tests +Please read the important note about setting `CUDA_VISIBLE_DEVICES` in [CT Perftest section](#ct-perftest). + **Sequential CT Perftest**: ```bash # Basic usage @@ -375,9 +380,12 @@ python main.py --debug sequential-ct-perftest ./config.yaml \ --json-output-path ./results.json # With Slurm -srun python main.py sequential-ct-perftest ./config.yaml \ +srun bash -c " + CUDA_VISIBLE_DEVICES=$SLURM_LOCALID + python main.py sequential-ct-perftest ./config.yaml \ --verify-buffers \ --json-output-path ./results.json +" ``` **CT Perftest**: diff --git a/benchmark/kvbench/commands/args.py b/benchmark/kvbench/commands/args.py index ca1790dde2..6898f64816 100644 --- a/benchmark/kvbench/commands/args.py +++ b/benchmark/kvbench/commands/args.py @@ -41,6 +41,18 @@ def cli_args(func): func = click.option("--num_requests", type=int, help="Number of requests")(func) func = click.option("--page_size", type=int, help="Page size")(func) func = click.option("--access_pattern", type=str, help="Access pattern")(func) + func = click.option( + "--source", + default="file", + type=str, + help="Source of the nixl descriptors [file, memory, gpu] (default: file)", + )(func) + func = click.option( + "--destination", + default="memory", + type=str, + help="Destination of the nixl descriptors [file, memory, gpu] (default: memory)", + )(func) return func @@ -57,22 +69,10 @@ def plan_args(func): def nixl_bench_args(func): """Decorator for NIXL benchmark arguments""" - func = click.option( - "--source", - default="file", - type=str, - help="Source of the nixl descriptors [file, memory, gpu] (default: file)", - )(func) - func = click.option( - "--destination", - default="memory", - type=str, - help="Destination of the nixl descriptors [file, memory, gpu] (default: memory)", - )(func) func = click.option( "--backend", type=str, - help="Communication backend [POSIX, GDS] (default: POSIX)", + help="Communication backend [UCX, UCX_MO, GDS, GDS_MT, POSIX, GPUNETIO, Mooncake, HF3FS, OBJ] (default: UCX)", )(func) func = click.option( "--worker_type", @@ -82,12 +82,12 @@ def nixl_bench_args(func): func = click.option( "--initiator_seg_type", type=str, - help="Memory segment type for initiator [DRAM, VRAM] (default: DRAM)", + help="Memory segment type for initiator [DRAM, VRAM, FILE, OBJ] (default: DRAM)", )(func) func = click.option( "--target_seg_type", type=str, - help="Memory segment type for target [DRAM, VRAM] (default: DRAM)", + help="Memory segment type for target [DRAM, VRAM, FILE, OBJ] (default: DRAM)", )(func) func = click.option( "--scheme", @@ -144,6 +144,9 @@ def nixl_bench_args(func): func = click.option("--enable_pt", is_flag=True, help="Enable progress thread")( func ) + func = click.option( + "--progress_threads", type=int, help="Number of progress threads (default: 0)" + )(func) func = click.option( "--device_list", type=str, help="Comma-separated device names (default: all)" )(func) @@ -151,7 +154,7 @@ def nixl_bench_args(func): "--runtime_type", type=str, help="Type of runtime to use [ETCD] (default: ETCD)" )(func) func = click.option( - "--etcd-endpoints", + "--etcd_endpoints", type=str, help="ETCD server URL for coordination (default: http://localhost:2379)", )(func) @@ -170,16 +173,97 @@ def nixl_bench_args(func): help="API type for POSIX operations [AIO, URING] (only used with POSIX backend", )(func) func = click.option( - "--enable-vmm", + "--enable_vmm", type=bool, help="Enable VMM memory allocation when DRAM is requested", )(func) func = click.option( - "--benchmark-group", + "--benchmark_group", type=str, help="Name of benchmark group (default: default). Use different names to run multiple benchmarks in parallel", default="default", )(func) + # Missing arguments from nixlbench + func = click.option( + "--num_files", + type=int, + help="Number of files used by benchmark (default: 1)", + )(func) + func = click.option( + "--large_blk_iter_ftr", + type=int, + help="Factor to reduce test iteration when testing large block size(>1MB) (default: 16)", + )(func) + func = click.option( + "--gds_batch_pool_size", + type=int, + help="Batch pool size for GDS operations (default: 32, only used with GDS backend)", + )(func) + func = click.option( + "--gds_batch_limit", + type=int, + help="Batch limit for GDS operations (default: 128, only used with GDS backend)", + )(func) + func = click.option( + "--gds_mt_num_threads", + type=int, + help="Number of threads used by GDS MT plugin (default: 1)", + )(func) + func = click.option( + "--gpunetio_device_list", + type=str, + help="Comma-separated GPU CUDA device id to use for communication (only used with GPUNETIO backend)", + )(func) + func = click.option( + "--hf3fs_iopool_size", + type=int, + help="Size of io memory pool for HF3FS backend (default: 64)", + )(func) + func = click.option( + "--obj_access_key", + type=str, + help="Access key for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_secret_key", + type=str, + help="Secret key for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_session_token", + type=str, + help="Session token for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_bucket_name", + type=str, + help="Bucket name for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_scheme", + type=str, + help="HTTP scheme for S3 backend [http, https] (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_region", + type=str, + help="Region for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_use_virtual_addressing", + type=bool, + help="Use virtual addressing for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_endpoint_override", + type=str, + help="Endpoint override for S3 backend (only used with OBJ backend)", + )(func) + func = click.option( + "--obj_req_checksum", + type=str, + help="Required checksum type for S3 backend [supported, required] (only used with OBJ backend)", + )(func) return func diff --git a/benchmark/kvbench/commands/nixlbench.py b/benchmark/kvbench/commands/nixlbench.py index 21d55d47f4..c4be865358 100644 --- a/benchmark/kvbench/commands/nixlbench.py +++ b/benchmark/kvbench/commands/nixlbench.py @@ -34,6 +34,7 @@ def __init__( check_consistency=False, device_list="all", enable_pt=False, + progress_threads=0, etcd_endpoints="http://localhost:2379", storage_enable_direct=False, filepath="", @@ -60,6 +61,20 @@ def __init__( warmup_iter=100, worker_type="nixl", benchmark_group="default", + gds_mt_num_threads=1, + gpunetio_device_list="0", + hf3fs_iopool_size=64, + obj_access_key="", + obj_secret_key="", + obj_session_token="", + obj_bucket_name="", + obj_scheme="http", + obj_region="eu-central-1", + obj_use_virtual_addressing=False, + obj_endpoint_override="", + obj_req_checksum="supported", + # Additional nixlbench arguments + large_blk_iter_ftr=16, ): """ Initialize a NIXLBench instance with benchmark configuration. @@ -72,6 +87,7 @@ def __init__( check_consistency (bool, optional): Whether to check consistency. Defaults to False. device_list (str, optional): List of devices to use. Defaults to "all". enable_pt (bool, optional): Whether to enable peer-to-peer transfer. Defaults to False. + progress_threads (int, optional): Number of progress threads (default: 0). etcd_endpoints (str, optional): ETCD endpoints for runtime. Defaults to "http://localhost:2379". storage_enable_direct (bool, optional): Whether to enable direct I/O for storage operations. Defaults to False. filepath (str, optional): Path for GDS and POSIX operations. Defaults to "". @@ -97,6 +113,19 @@ def __init__( total_buffer_size (int, optional): Total buffer size. Defaults to 8589934592. warmup_iter (int, optional): Number of warmup iterations. Defaults to 100. worker_type (str, optional): Type of worker. Defaults to "nixl". + gds_mt_num_threads (int, optional): Number of threads for GDS_MT plugin. Defaults to 1. + gpunetio_device_list (str, optional): GPU device list for GPUNETIO plugin. Defaults to "0". + hf3fs_iopool_size (int, optional): IO pool size for HF3FS plugin. Defaults to 64. + obj_access_key (str, optional): Access key for OBJ/S3 plugin. Defaults to "". + obj_secret_key (str, optional): Secret key for OBJ/S3 plugin. Defaults to "". + obj_session_token (str, optional): Session token for OBJ/S3 plugin. Defaults to "". + obj_bucket_name (str, optional): Bucket name for OBJ/S3 plugin. Defaults to "". + obj_scheme (str, optional): HTTP scheme for OBJ/S3 plugin. Defaults to "http". + obj_region (str, optional): Region for OBJ/S3 plugin. Defaults to "eu-central-1". + obj_use_virtual_addressing (bool, optional): Use virtual addressing for OBJ/S3. Defaults to False. + obj_endpoint_override (str, optional): Endpoint override for OBJ/S3. Defaults to "". + obj_req_checksum (str, optional): Required checksum for OBJ/S3. Defaults to "supported". + large_blk_iter_ftr (int, optional): Factor to reduce iterations for large blocks. Defaults to 16. """ self.model = model self.model_config = model_config @@ -105,6 +134,7 @@ def __init__( self.check_consistency = check_consistency self.device_list = device_list self.enable_pt = enable_pt + self.progress_threads = progress_threads self.etcd_endpoints = etcd_endpoints self.storage_enable_direct = storage_enable_direct self.filepath = filepath @@ -130,6 +160,19 @@ def __init__( self.total_buffer_size = total_buffer_size self.warmup_iter = warmup_iter self.worker_type = worker_type + self.gds_mt_num_threads = gds_mt_num_threads + self.gpunetio_device_list = gpunetio_device_list + self.hf3fs_iopool_size = hf3fs_iopool_size + self.obj_access_key = obj_access_key + self.obj_secret_key = obj_secret_key + self.obj_session_token = obj_session_token + self.obj_bucket_name = obj_bucket_name + self.obj_scheme = obj_scheme + self.obj_region = obj_region + self.obj_use_virtual_addressing = obj_use_virtual_addressing + self.obj_endpoint_override = obj_endpoint_override + self.obj_req_checksum = obj_req_checksum + self.large_blk_iter_ftr = large_blk_iter_ftr self._override_defaults() def set_io_size(self, io_size: int): @@ -137,18 +180,18 @@ def set_io_size(self, io_size: int): self.max_block_size = io_size def _configure_gds(self, source: str, destination: str): + """Configure GDS and GDS_MT plugins (same logic for both)""" if source == "file": - # this is a READ from GDS to GPU self.op_type = "READ" self.target_seg_type = "VRAM" elif source == "gpu": - # this is a WRITE from GPU to GDS self.op_type = "WRITE" - self.target_seg_type = "VRAM" + self.target_seg_type = "FILE" else: - raise ValueError(f"Invalid source for GDS: {source}") + raise ValueError(f"Invalid source for GDS/GDS_MT: {source}") def _configure_posix(self, source: str, destination: str): + """Configure POSIX and HF3FS plugins (same logic for both)""" if source == "file": self.op_type = "READ" self.target_seg_type = "DRAM" @@ -156,34 +199,52 @@ def _configure_posix(self, source: str, destination: str): self.op_type = "WRITE" self.initiator_seg_type = "DRAM" else: - raise ValueError(f"Invalid source for POSIX: {source}") + raise ValueError(f"Invalid source for POSIX/HF3FS: {source}") + + def _configure_ucx(self, backend: str, source: str, destination: str): + """Configure UCX, UCX_MO, GPUNETIO, and Mooncake plugins (same logic for all)""" + arg_to_seg_type = { + "memory": "DRAM", + "gpu": "VRAM", + } + + backend = backend.upper() + try: + self.initiator_seg_type = arg_to_seg_type[source] + except KeyError: + raise ValueError( + f"Invalid source for {backend}: {source}, valid sources are: {arg_to_seg_type.keys()}" + ) + try: + self.target_seg_type = arg_to_seg_type[destination] + except KeyError: + raise ValueError( + f"Invalid destination for {backend}: {destination}, valid destinations are: {arg_to_seg_type.keys()}" + ) + + def _configure_obj(self, source: str, destination: str): + """Configure OBJ plugin for object storage operations""" + if source == "memory": + self.target_seg_type = "OBJ" + elif destination == "memory": + self.initiator_seg_type = "OBJ" + else: + raise ValueError(f"Invalid source for OBJ: {source}") def configure_segment_type(self, backend: str, source: str, destination: str): - if backend.lower() == "gds": + backend_lower = backend.lower() + + if backend_lower in ["gds", "gds_mt"]: self._configure_gds(source, destination) - elif backend.lower() == "posix": + elif backend_lower in ["posix", "hf3fs"]: self._configure_posix(source, destination) + elif backend_lower in ["ucx", "ucx_mo", "gpunetio", "mooncake"]: + self._configure_ucx(backend_lower, source, destination) + elif backend_lower == "obj": + self._configure_obj(source, destination) else: raise ValueError(f"Invalid backend: {backend}") - # if backend == "GDS" or backend == "POSIX": - # if source == "file": - # # this is a READ from GDS to GPU - # self.op_type = "READ" - # self.target_seg_type = "VRAM" - # elif source == "gpu": - # # this is a WRITE from GPU to GDS - # self.op_type = "WRITE" - # self.target_seg_type = "VRAM" - - # elif source == "memory": - # # this is a WRITE from memory to GDS - # self.op_type = "WRITE" - # self.initiator_seg_type = "DRAM" - # self.target_seg_type = "DRAM" - # else: - # raise ValueError(f"Invalid backend: {backend}") - def configure_scheme(self, scheme: str = "pairwise", direction: str = "isl"): """ Configure the scheme based on the model configuration. @@ -231,6 +292,7 @@ def _params(self): "check_consistency": self.check_consistency, "device_list": self.device_list, "enable_pt": self.enable_pt, + "progress_threads": self.progress_threads, "etcd_endpoints": self.etcd_endpoints, "storage_enable_direct": self.storage_enable_direct, "filepath": self.filepath, @@ -256,6 +318,20 @@ def _params(self): "total_buffer_size": self.total_buffer_size, "warmup_iter": self.warmup_iter, "worker_type": self.worker_type, + "gds_mt_num_threads": self.gds_mt_num_threads, + "gpunetio_device_list": self.gpunetio_device_list, + "hf3fs_iopool_size": self.hf3fs_iopool_size, + "obj_access_key": self.obj_access_key, + "obj_secret_key": self.obj_secret_key, + "obj_session_token": self.obj_session_token, + "obj_bucket_name": self.obj_bucket_name, + "obj_scheme": self.obj_scheme, + "obj_region": self.obj_region, + "obj_use_virtual_addressing": self.obj_use_virtual_addressing, + "obj_endpoint_override": self.obj_endpoint_override, + "obj_req_checksum": self.obj_req_checksum, + # Additional nixlbench parameters + "large_blk_iter_ftr": self.large_blk_iter_ftr, } @staticmethod @@ -274,6 +350,7 @@ def defaults(): "check_consistency": False, "device_list": "all", "enable_pt": False, + "progress_threads": 0, "etcd_endpoints": "http://localhost:2379", "storage_enable_direct": False, "filepath": "", @@ -300,6 +377,20 @@ def defaults(): "warmup_iter": 100, "worker_type": "nixl", "benchmark_group": "default", + "gds_mt_num_threads": 1, + "gpunetio_device_list": "0", + "hf3fs_iopool_size": 64, + "obj_access_key": "", + "obj_secret_key": "", + "obj_session_token": "", + "obj_bucket_name": "", + "obj_scheme": "http", + "obj_region": "eu-central-1", + "obj_use_virtual_addressing": False, + "obj_endpoint_override": "", + "obj_req_checksum": "supported", + # Additional nixlbench defaults + "large_blk_iter_ftr": 16, } def plan(self, format: str = "text"): @@ -333,7 +424,6 @@ def should_include(name, value, include_defaults=False): for key, value in params.items(): if value is not None: merged_params[key] = value - # print(json.dumps(merged_params)) return merged_params else: # for text format, exclude defaults to keep command concise for name, value in params.items(): diff --git a/benchmark/kvbench/main.py b/benchmark/kvbench/main.py index 6ca138fb74..0b1f922dba 100644 --- a/benchmark/kvbench/main.py +++ b/benchmark/kvbench/main.py @@ -76,10 +76,10 @@ def cli(debug): @cli.command("plan") +@cli_args @common_args @plan_args @nixl_bench_args -@cli_args def plan_command(model, model_config, model_configs, format, **kwargs): """Display the recommended configuration for nixlbench""" if not model: @@ -213,9 +213,9 @@ def plan_command(model, model_config, model_configs, format, **kwargs): @cli.command("profile") +@cli_args @common_args @nixl_bench_args -@cli_args def profile_command(model, model_config, **kwargs): """Run nixlbench""" if not model or not model_config: @@ -254,8 +254,8 @@ def profile_command(model, model_config, **kwargs): @cli.command("kvcache") -@common_args @cli_args +@common_args def kvcache_command(model, model_config, **kwargs): """Display kvcache information""" if not model or not model_config: diff --git a/benchmark/kvbench/models/model_config.py b/benchmark/kvbench/models/model_config.py index c3b3fb66df..540cddcfd8 100644 --- a/benchmark/kvbench/models/model_config.py +++ b/benchmark/kvbench/models/model_config.py @@ -20,6 +20,10 @@ import yaml # type: ignore +from nixl.logging import get_logger + +logger = get_logger(__name__) + @dataclass class StrategyConfig: @@ -121,7 +125,7 @@ def from_yaml_files(cls, yaml_paths: List[str]) -> "ModelConfig": config_dict = yaml.safe_load(f) config = config.update(config_dict) else: - print(f"Warning: Config file not found: {path}") + logger.warning("Config file not found: %s", path) return config diff --git a/benchmark/kvbench/pyproject.toml b/benchmark/kvbench/pyproject.toml index 6091f6d349..c30e4ce1cb 100644 --- a/benchmark/kvbench/pyproject.toml +++ b/benchmark/kvbench/pyproject.toml @@ -21,13 +21,13 @@ readme = "README.md" requires-python = ">=3.9" dependencies = [ "click==8.1.7", - "numpy==2.3.1", "pytest==7.4.4", "pyyaml>=6.0.2", "etcd3>=0.12.0", "tabulate==0.9.0", - "torch==2.7.0", + "torch>=2.7.0", "tqdm==4.66.5", + "numpy", "nixl", ] diff --git a/benchmark/kvbench/runtime/etcd_rt.py b/benchmark/kvbench/runtime/etcd_rt.py index 4f3adecdbc..da70c4cb2f 100644 --- a/benchmark/kvbench/runtime/etcd_rt.py +++ b/benchmark/kvbench/runtime/etcd_rt.py @@ -13,7 +13,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import logging import os import pickle import re @@ -23,9 +22,11 @@ import etcd3 +from nixl.logging import get_logger + from .rt_base import ReduceOp, _RTUtils -log = logging.getLogger(__name__) +logger = get_logger(__name__) def int_to_bytes(val: int) -> bytes: @@ -75,7 +76,7 @@ def __init__( f"Invalid etcd endpoint format: {etcd_endpoints}, expected format is [http://]host[:port]" ) - log.info(f"ETCD client initialized with host {host} & port {port}") + logger.info("ETCD client initialized with host %s & port %d", host, port) try: self.client = etcd3.client(host=host, port=port) @@ -83,7 +84,7 @@ def __init__( raise ValueError(f"Failed to initialize ETCD client: {e}") if self.rank == 0: - log.info(f"Wiping ETCD prefix {self.prefix}") + logger.info("Wiping ETCD prefix %s", self.prefix) self.client.delete_prefix(self.prefix) def destroy_dist(self): @@ -124,7 +125,13 @@ def barrier(self, ranks: Optional[List[int]] = None, timeout_sec=600): ): if timeout_sec and time.time() - start_time > timeout_sec: raise TimeoutError( - f"[Rank {self.rank}] ROOT - Barrier {key} timed out after {timeout_sec} seconds, current value: {self.client.get(key)}, waiting for val={len(ranks)} (i.e all the ranks have entered the barrier), (ranks: {ranks})" + "[Rank %d] ROOT - Barrier %s timed out after %.3f seconds, current value: %s, waiting for val=%d (i.e all the ranks have entered the barrier), (ranks: %s)", + self.rank, + key, + timeout_sec, + self.client.get(key), + len(ranks), + ranks, ) else: my_index = ranks.index(self.rank) @@ -207,7 +214,7 @@ def all_reduce( val = self.client.get(f"{self.prefix}/all_reduce/{dest_rank}")[0] vals.append(pickle.loads(val)) - print(vals) + logger.debug("All reduce values: %s", vals) if op == ReduceOp.SUM: final_val = [sum(col) for col in zip(*vals)] elif op == ReduceOp.AVG: @@ -235,7 +242,7 @@ def _get_group_id(self, ranks: List[int]) -> int: if not os.environ.get("NIXL_ETCD_NAMESPACE"): - log.warning( + logger.warning( "Environment variable NIXL_ETCD_NAMESPACE is not set, using default prefix /nixl/kvbench. " "Note that it can lead to conflicts if multiple instances of KVBench are running. " "To avoid this, set NIXL_ETCD_NAMESPACE to a unique value for each instance of KVBench. " diff --git a/benchmark/kvbench/test/custom_traffic_perftest.py b/benchmark/kvbench/test/custom_traffic_perftest.py index 6a3bcf2cf1..51b4be2958 100644 --- a/benchmark/kvbench/test/custom_traffic_perftest.py +++ b/benchmark/kvbench/test/custom_traffic_perftest.py @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -import logging +import os import time from test.traffic_pattern import TrafficPattern from typing import Literal, Optional, Tuple @@ -24,8 +24,9 @@ from tabulate import tabulate from nixl._api import nixl_agent +from nixl.logging import get_logger -log = logging.getLogger(__name__) +logger = get_logger(__name__) class NixlHandle: @@ -62,13 +63,20 @@ def __init__( if shards > 1: raise ValueError("Sharding is not supported yet") - log.debug( - f"[Rank {dist_rt.get_rank()}] Initializing NixlBuffer with size {size}, device {device}, shards {shards}, fill_value {fill_value}" + logger.debug( + "[Rank %d] Initializing NixlBuffer with size %d, device %s, shards %d, fill_value %d", + dist_rt.get_rank(), + size, + device, + shards, + fill_value, ) self.buf = torch.full((size,), fill_value, dtype=dtype, device=device) - log.debug( - f"[Rank {dist_rt.get_rank()}] Registering memory for buffer {self.buf}" + logger.debug( + "[Rank %d] Registering memory for buffer %s", + dist_rt.get_rank(), + self.buf, ) self.reg_descs = nixl_agent.get_reg_descs(self.buf) assert ( @@ -107,6 +115,14 @@ def __init__( self.nixl_agent = nixl_agent(f"{self.my_rank}") + if ( + not os.environ.get("CUDA_VISIBLE_DEVICES") + and self.traffic_pattern.mem_type == "cuda" + ): + logger.warning( + "Cuda buffers detected, but the env var CUDA_VISIBLE_DEVICES is not set, this will cause every process in the same host to use the same GPU device." + ) + """Initialize the buffers, one big send and recv buffer is used for all the transfers it has to be chunked inside each transfer to get buffers per ranks the buffer is big enough to handle any of the transfers @@ -137,7 +153,7 @@ def _barrier_tp(self, tp: TrafficPattern, senders_only=True): def _share_md(self) -> None: """Share agent metadata between all ranks. (Need to be run after registering buffers)""" - log.debug(f"[Rank {self.my_rank}] Sharing MD") + logger.debug("[Rank %d] Sharing MD", self.my_rank) md = self.nixl_agent.get_agent_metadata() mds = dist_rt.allgather_obj(md) for other_rank, metadata in enumerate(mds): @@ -202,7 +218,7 @@ def _prepare_tp( send_bufs, recv_bufs = self._get_bufs(tp) - log.debug(f"[Rank {self.my_rank}] Sharing recv buf descs") + logger.debug("[Rank %d] Sharing recv buf descs", self.my_rank) dst_bufs_descs = self._share_recv_buf_descs(recv_bufs) handles: list[NixlHandle] = [] @@ -212,8 +228,12 @@ def _prepare_tp( xfer_desc = self.nixl_agent.get_xfer_descs(buf) - log.debug( - f"[Rank {self.my_rank}] Initializing xfer for {other} - xfer desc: {xfer_desc}, dst buf desc: {dst_bufs_descs[other]}" + logger.debug( + "[Rank %d] Initializing xfer for %d - xfer desc: %s, dst buf desc: %s", + self.my_rank, + other, + xfer_desc, + dst_bufs_descs[other], ) handle = self.nixl_agent.initialize_xfer( "WRITE", @@ -268,11 +288,11 @@ def _wait(self, handles: list[NixlHandle]): handles = pending def _destroy(self, handles: list[NixlHandle]): - log.debug(f"[Rank {self.my_rank}] Releasing XFER handles") + logger.debug("[Rank %d] Releasing XFER handles", self.my_rank) for handle in handles: self.nixl_agent.release_xfer_handle(handle.handle) - log.debug(f"[Rank {self.my_rank}] Removing remote agents") + logger.debug("[Rank %d] Removing remote agents", self.my_rank) for other_rank in range(self.world_size): if other_rank == self.my_rank: continue @@ -281,7 +301,7 @@ def _destroy(self, handles: list[NixlHandle]): self._destroy_buffers() def _destroy_buffers(self): - log.debug(f"[Rank {self.my_rank}] Destroying buffers") + logger.debug("[Rank %d] Destroying buffers", self.my_rank) self.send_buf.destroy() self.recv_buf.destroy() @@ -301,14 +321,14 @@ def _verify_tp( for r, recv_buf in enumerate(recv_bufs): if recv_buf is None: if tp.matrix[r][self.my_rank] > 0: - log.error( + logger.error( f"Rank {self.my_rank} expected {tp.matrix[r][self.my_rank]} bytes from rank {r}, but got 0" ) raise RuntimeError("Buffer verification failed") continue if print_recv_buffers: - log.info(f"Recv buffer {r}:\n{recv_buf.buf}") + logger.info("Recv buffer %d:\n%s", r, recv_buf.buf) # recv_buf has to be filled with the rank of the sender # and its size has to be the same as matrix[r][my_rank] @@ -337,7 +357,7 @@ def run( Returns: Total execution time in seconds """ - log.debug(f"[Rank {self.my_rank}] Running CT perftest") + logger.debug("[Rank %d] Running CT perftest", self.my_rank) self._share_md() handles, send_bufs, recv_bufs = self._prepare_tp(self.traffic_pattern) @@ -377,7 +397,10 @@ def run( total_size_gb, ] ] - print(tabulate(data, headers=headers, floatfmt=".6f")) + logger.info( + "Performance metrics:\n%s", + tabulate(data, headers=headers, floatfmt=".6f"), + ) if verify_buffers: self._verify_tp(self.traffic_pattern, recv_bufs, print_recv_buffers) diff --git a/benchmark/kvbench/test/inference_workload_matgen.py b/benchmark/kvbench/test/inference_workload_matgen.py index 42d42a0ebe..dd23647fb7 100644 --- a/benchmark/kvbench/test/inference_workload_matgen.py +++ b/benchmark/kvbench/test/inference_workload_matgen.py @@ -57,6 +57,10 @@ import yaml from tqdm import tqdm +from nixl.logging import get_logger + +logger = get_logger(__name__) + @dataclass class ModelConfig: @@ -190,7 +194,7 @@ def gen_batches( curr_mem = 0 if curr: batches.append(Batch(user_requests=curr)) - print(f"Last batch is incomplete, his size is {len(curr)}") + logger.warning("Last batch is incomplete, with size %d", len(curr)) return batches @@ -268,8 +272,6 @@ def gen_matrix( num_peers = int(num_peers) buf_size = kv_slice_size / num_peers - # print(f"kv_size: {format_size(kv_size)}, kv_slice_size: {format_size(kv_slice_size)}, buf_size: {format_size(buf_size)}, num_peers: {num_peers}") - mat = np.zeros((world_size, world_size)) dst_iter = iter(decode_worker) @@ -389,11 +391,11 @@ def main( decode_workers = reordered - print(f"Prefill workers: {prefill_workers}") - print(f"Decode workers: {decode_workers}") + logger.info("Prefill workers: %s", prefill_workers) + logger.info("Decode workers: %s", decode_workers) batches = gen_batches(num_user_requests, task_config, model_config) - print(f"Generated {len(batches)} batches") + logger.info("Generated %d batches", len(batches)) matrices = gen_matrices_and_compute_time( batches, prefill_workers, @@ -406,7 +408,7 @@ def main( # Save matrices and metadata to files results_dir = results_dir or Path(f"matrices_{world_size}ranks") results_dir = Path(results_dir) - print(f"Saving {len(matrices)} matrices to {results_dir}") + logger.info("Saving %d matrices to %s", len(matrices), results_dir) results_dir.mkdir(parents=True, exist_ok=True) metadata: dict[str, Any] = { @@ -435,7 +437,7 @@ def main( metadata_path = results_dir / "metadata.yaml" with open(metadata_path, "w") as f: yaml.dump(metadata, f) - print(f"Saved metadata to {metadata_path}") + logger.info("Saved metadata to %s", metadata_path) if __name__ == "__main__": @@ -603,10 +605,6 @@ def generate( max_batch_mem=max_batch_mem, ) - # world_size = num_prefill_nodes * prefill_tp + num_decode_nodes * decode_tp - # print(f"World size: {world_size}") - # print(f"Model config: {model_config}") - main( num_user_requests=num_user_requests, task_config=task_config, diff --git a/benchmark/kvbench/test/sequential_custom_traffic_perftest.py b/benchmark/kvbench/test/sequential_custom_traffic_perftest.py index 278ac7476c..c68a5e6c98 100644 --- a/benchmark/kvbench/test/sequential_custom_traffic_perftest.py +++ b/benchmark/kvbench/test/sequential_custom_traffic_perftest.py @@ -16,7 +16,7 @@ """Sequential is different from multi in that every rank processes only one TP at a time, but they can process different ones""" import json -import logging +import os import time from collections import defaultdict from itertools import chain @@ -29,8 +29,9 @@ from tabulate import tabulate from nixl._api import nixl_agent +from nixl.logging import get_logger -log = logging.getLogger(__name__) +logger = get_logger(__name__) class SequentialCTPerftest(CTPerftest): @@ -59,11 +60,17 @@ def __init__( self.n_isolation_iters = n_isolation_iters self.warmup_iters = warmup_iters - log.debug(f"[Rank {self.my_rank}] Initializing Nixl agent") + logger.debug("[Rank %d] Initializing Nixl agent", self.my_rank) self.nixl_agent = nixl_agent(f"{self.my_rank}") for tp in self.traffic_patterns: self._check_tp_config(tp) + if not os.environ.get("CUDA_VISIBLE_DEVICES") and any( + tp.mem_type == "cuda" for tp in self.traffic_patterns + ): + logger.warning( + "Cuda buffers detected, but the env var CUDA_VISIBLE_DEVICES is not set, this will cause every process in the same host to use the same GPU device." + ) assert "UCX" in self.nixl_agent.get_plugin_list(), "UCX plugin is not loaded" # NixlBuffer caches buffers and reuse them if they are big enough, let's initialize them once, with the largest needed size @@ -71,7 +78,7 @@ def __init__( self.recv_buf_by_mem_type: dict[str, NixlBuffer] = {} def _init_buffers(self): - log.debug(f"[Rank {self.my_rank}] Initializing buffers") + logger.debug("[Rank %d] Initializing buffers", self.my_rank) max_src_by_mem_type = defaultdict(int) max_dst_by_mem_type = defaultdict(int) @@ -98,14 +105,14 @@ def _init_buffers(self): ) def _destroy_buffers(self): - log.debug(f"[Rank {self.my_rank}] Destroying buffers") + logger.debug("[Rank %d] Destroying buffers", self.my_rank) for buf in chain( self.send_buf_by_mem_type.values(), self.recv_buf_by_mem_type.values() ): buf.destroy() def _get_bufs(self, tp: TrafficPattern): - log.debug(f"[Rank {self.my_rank}] Getting buffers for TP {tp.id}") + logger.debug("[Rank %d] Getting buffers for TP %s", self.my_rank, tp.id) send_bufs = [None for _ in range(self.world_size)] recv_bufs = [None for _ in range(self.world_size)] @@ -152,7 +159,7 @@ def run( This method initializes and executes multiple traffic patterns simultaneously, measures their performance, and optionally verifies the results. """ - log.debug(f"[Rank {self.my_rank}] Running sequential CT perftest") + logger.debug("[Rank %d] Running sequential CT perftest", self.my_rank) self._init_buffers() self._share_md() @@ -168,7 +175,7 @@ def run( tp_bufs = [] s = time.time() - log.info(f"[Rank {self.my_rank}] Preparing TPs") + logger.info("[Rank %d] Preparing TPs", self.my_rank) for i, tp in enumerate(self.traffic_patterns): handles, send_bufs, recv_bufs = self._prepare_tp(tp) tp_bufs.append((send_bufs, recv_bufs)) @@ -190,8 +197,9 @@ def run( dist_rt.barrier() # Isolated mode - Measure SOL for every matrix - log.info( - f"[Rank {self.my_rank}] Running isolated benchmark (to measure perf without noise)" + logger.info( + "[Rank %d] Running isolated benchmark (to measure perf without noise)", + self.my_rank, ) my_isolated_tp_latencies: list[float] = [0 for _ in tp_handles] @@ -211,8 +219,13 @@ def run( my_isolated_tp_latencies[tp_ix] += e - t self._barrier_tp(tp) - log.debug( - f"[Rank {self.my_rank}] Ran {self.n_isolation_iters} isolated iters for tp {tp_ix}/{len(tp_handles)}, took {e - t} secs" + logger.debug( + "[Rank %d] Ran %d isolated iters for tp %d/%d, took %.3f secs", + self.my_rank, + self.n_isolation_iters, + tp_ix, + len(tp_handles), + e - t, ) my_isolated_tp_latencies[tp_ix] /= self.n_isolation_iters @@ -230,18 +243,21 @@ def run( if tp_lats: isolated_tp_latencies_ms.append(max(tp_lats) * 1e3) - log.info(f"[Rank {self.my_rank}] Running workload benchmark") + logger.info("[Rank %d] Running workload benchmark", self.my_rank) # Workload mode - Measure perf of the matrices while running the full workload for iter_ix in range(self.n_iters): - log.debug( - f"[Rank {self.my_rank}] Running iteration {iter_ix + 1}/{self.n_iters}" + logger.debug( + "[Rank %d] Running iteration %d/%d", + self.my_rank, + iter_ix + 1, + self.n_iters, ) iter_metadata = results["metadata"]["iters"][iter_ix] tp_starts: list[float | None] = [None] * len(tp_handles) tp_ends: list[float | None] = [None] * len(tp_handles) - log.debug(f"[Rank {self.my_rank}] Warmup done.") + logger.debug("[Rank %d] Warmup done.", self.my_rank) dist_rt.barrier(timeout_sec=None) iter_metadata["start_ts"] = time.time() @@ -256,14 +272,22 @@ def run( time.sleep(tp.sleep_before_launch_sec) # Run TP - log.debug(f"[Rank {self.my_rank}] Running TP {tp_ix}/{len(tp_handles)}") + logger.debug( + "[Rank %d] Running TP %d/%d", + self.my_rank, + tp_ix, + len(tp_handles), + ) tp_start_ts = time.time() self._run_tp(handles, blocking=True) tp_end_ts = time.time() - log.debug( - f"[Rank {self.my_rank}] TP {tp_ix} took {tp_end_ts - tp_start_ts} seconds" + logger.debug( + "[Rank %d] TP %d took %.3f seconds", + self.my_rank, + tp_ix, + tp_end_ts - tp_start_ts, ) tp_starts[tp_ix] = tp_start_ts @@ -300,12 +324,27 @@ def run( else: tp_latencies_ms.append((max(ends) - min(starts)) * 1e3) + mean_bw = 0.0 + for rank in tp.senders_ranks(): + rank_start = tp_starts_by_ranks[rank][i] + rank_end = tp_ends_by_ranks[rank][i] + if not rank_start or not rank_end: + raise ValueError( + f"Rank {rank} has no start or end time, but participated in TP, this is not normal." + ) + mean_bw += ( + tp.total_src_size(rank) * 1e-9 / (rank_end - rank_start) + ) + + mean_bw /= len(tp.senders_ranks()) + if self.my_rank == 0: headers = [ "Transfer size (GB)", "Latency (ms)", "Isolated Latency (ms)", "Num Senders", + "Mean BW (GB/s)", # Bandwidth ] data = [ [ @@ -313,12 +352,12 @@ def run( tp_latencies_ms[i], isolated_tp_latencies_ms[i], len(tp.senders_ranks()), + mean_bw, ] for i, tp in enumerate(self.traffic_patterns) ] - print( - f"Iteration {iter_ix + 1}/{self.n_iters}\n", - tabulate(data, headers=headers, floatfmt=".3f"), + logger.info( + f"Iteration {iter_ix + 1}/{self.n_iters}\n{tabulate(data, headers=headers, floatfmt='.3f')}" ) if verify_buffers: @@ -332,6 +371,7 @@ def run( "latency": tp_latencies_ms[i], "isolated_latency": isolated_tp_latencies_ms[i], "num_senders": len(tp.senders_ranks()), + "mean_bw": mean_bw, "min_start_ts": min( filter( None, @@ -357,12 +397,12 @@ def run( results["metadata"]["finished_ts"] = time.time() if json_output_path and self.my_rank == 0: - log.info(f"Saving results to {json_output_path}") + logger.info("Saving results to %s", json_output_path) with open(json_output_path, "w") as f: json.dump(results, f) # Destroy - log.info(f"[Rank {self.my_rank}] Finished run, destroying objects") + logger.info("[Rank %d] Finished run, destroying objects", self.my_rank) self._destroy(handles) def _write_yaml_results( @@ -411,6 +451,6 @@ def _write_yaml_results( try: with open(output_path, "w") as f: yaml.dump(results, f, default_flow_style=False, sort_keys=False) - log.info(f"Results saved to YAML file: {output_path}") + logger.info("Results saved to YAML file: %s", output_path) except Exception as e: - log.error(f"Failed to write YAML results to {output_path}: {e}") + logger.error("Failed to write YAML results to %s: %s", output_path, e) diff --git a/benchmark/kvbench/test/traffic_pattern.py b/benchmark/kvbench/test/traffic_pattern.py index e771e58551..1f612279af 100644 --- a/benchmark/kvbench/test/traffic_pattern.py +++ b/benchmark/kvbench/test/traffic_pattern.py @@ -12,15 +12,12 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -import logging from dataclasses import dataclass, field from typing import ClassVar, Literal, Optional import numpy as np import torch -log = logging.getLogger(__name__) - @dataclass class TrafficPattern: diff --git a/benchmark/nixlbench/README.md b/benchmark/nixlbench/README.md index a0f8dc4c1d..eef45d1f15 100644 --- a/benchmark/nixlbench/README.md +++ b/benchmark/nixlbench/README.md @@ -21,19 +21,41 @@ A benchmarking tool for the NVIDIA Inference Xfer Library (NIXL) that uses ETCD ## Features -- Benchmarks NIXL performance across different backends -- Supports multiple communication patterns +- Benchmarks NIXL performance across multiple backends: + - **Network backends**: UCX, UCX_MO, GPUNETIO, Mooncake + - **Storage backends**: GDS, GDS_MT, POSIX, HF3FS, OBJ (S3) +- Supports multiple communication patterns: + - **Pairwise**: Point-to-point communication between pairs + - **Many-to-one**: Multiple initiators to single target + - **One-to-many**: Single initiator to multiple targets + - **TP (Tensor Parallel)**: Optimized for distributed training workloads - Tests both CPU (DRAM) and GPU (VRAM) memory transfers +- Support for multiple worker types: + - **NIXL worker**: Full-featured with all backend support + - **NVSHMEM worker**: GPU-focused with VRAM-only transfers - Uses ETCD for worker coordination - ideal for containerized and cloud-native environments +- Multi-threading support with configurable progress threads +- VMM memory allocation support for CUDA Fabric +- Comprehensive performance metrics with latency percentiles +- Data consistency validation for reliability testing ## Building ### Prerequisites -- NIXL Library -- CUDA Toolkit -- GFlags -- ETCD C++ client (https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3) +#### Required Dependencies +- **NIXL Library** - NVIDIA Inference Xfer Library +- **GFlags** - Command line flag processing +- **OpenMP** - Multi-threading support +- **ETCD C++ client** - Coordination runtime (https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3) + +#### Optional Dependencies +- **CUDA Toolkit** - Required for VRAM operations and GPU backends +- **NVSHMEM** - Required for NVSHMEM worker type +- **Backend-specific libraries** as needed: + - UCX for network communication + - GDS/cuFile for GPU Direct Storage + - io_uring for POSIX URING operations ### Building with Meson @@ -52,53 +74,134 @@ meson install #### Custom Dependency Paths -If NIXL is installed in a non-standard location, you can specify the path: +If dependencies are installed in non-standard locations, you can specify their paths: ```bash -# With custom NIXL path -meson setup /path/to/build/dir -Dnixl_path=/path/to/nixl/installation - -# To view all meson project options -meson configure /path/to/build/dir +# With custom dependency paths +meson setup build \ + -Dnixl_path=/path/to/nixl/installation \ + -Dcudapath_inc=/path/to/cuda/include \ + -Dcudapath_lib=/path/to/cuda/lib64 \ + -Detcd_inc_path=/path/to/etcd/include \ + -Detcd_lib_path=/path/to/etcd/lib \ + -Dnvshmem_inc_path=/path/to/nvshmem/include \ + -Dnvshmem_lib_path=/path/to/nvshmem/lib + +# To view all available meson project options +meson configure build ``` +#### Available Build Options +- `nixl_path`: Path to NIXL installation (default: /usr/local) +- `cudapath_inc`: Include path for CUDA +- `cudapath_lib`: Library path for CUDA +- `cudapath_stub`: Extra stub path for CUDA +- `etcd_inc_path`: Path to ETCD C++ client includes +- `etcd_lib_path`: Path to ETCD C++ client library +- `nvshmem_inc_path`: Path to NVSHMEM include directory +- `nvshmem_lib_path`: Path to NVSHMEM library directory + ## Usage ### Basic Usage ```bash -# Run the benchmark with default settings -# using ETCD runtime and ETCD server -./nixlbench --etcd-endpoints http://etcd-server:2379 --backend UCX --initiator_seg_type VRAM +# Run basic UCX benchmark with VRAM transfers +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX --initiator_seg_type VRAM --target_seg_type VRAM + +# Run storage benchmark with GDS backend +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend GDS --filepath /mnt/storage/testfile + +# Run S3 object storage benchmark +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend OBJ --obj_bucket_name my-bucket --obj_access_key $AWS_ACCESS_KEY_ID --obj_secret_key $AWS_SECRET_ACCESS_KEY + +# Run multi-threaded benchmark with progress threads +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX --num_threads 4 --enable_pt --progress_threads 2 ``` ### Command Line Options +#### Core Options ``` ---backend NAME # Communication backend [UCX, UCX_MO] (default: UCX) +--runtime_type NAME # Type of runtime to use [ETCD] (default: ETCD) --worker_type NAME # Worker to use to transfer data [nixl, nvshmem] (default: nixl) +--backend NAME # Communication backend [UCX, UCX_MO, GDS, GDS_MT, POSIX, GPUNETIO, Mooncake, HF3FS, OBJ] (default: UCX) +--benchmark_group NAME # Name of benchmark group for parallel runs (default: default) +``` + +#### Memory and Transfer Configuration +``` --initiator_seg_type TYPE # Memory segment type for initiator [DRAM, VRAM] (default: DRAM) --target_seg_type TYPE # Memory segment type for target [DRAM, VRAM] (default: DRAM) --scheme NAME # Communication scheme [pairwise, manytoone, onetomany, tp] (default: pairwise) --mode MODE # Process mode [SG (Single GPU per proc), MG (Multi GPU per proc)] (default: SG) --op_type TYPE # Operation type [READ, WRITE] (default: WRITE) --check_consistency # Enable consistency checking ---total_buffer_size SIZE # Total buffer size (default: 8GiB) +--total_buffer_size SIZE # Total buffer size across devices per process (default: 8GiB) --start_block_size SIZE # Starting block size (default: 4KiB) --max_block_size SIZE # Maximum block size (default: 64MiB) --start_batch_size SIZE # Starting batch size (default: 1) --max_batch_size SIZE # Maximum batch size (default: 1) +``` + +#### Performance and Threading +``` --num_iter NUM # Number of iterations (default: 1000) --warmup_iter NUM # Number of warmup iterations (default: 100) +--large_blk_iter_ftr NUM # Factor to reduce transfer iteration for block size above 1MB (default: 16) --num_threads NUM # Number of threads used by benchmark (default: 1) --num_initiator_dev NUM # Number of devices in initiator processes (default: 1) --num_target_dev NUM # Number of devices in target processes (default: 1) ---enable_pt # Enable progress thread ---device_list LIST # Comma-separated device names (default: all) ---runtime_type NAME # Type of runtime to use [ETCD] (default: ETCD) ---etcd-endpoints URL # ETCD server URL for coordination (default: http://localhost:2379) +--enable_pt # Enable progress thread (only used with nixl worker) +--progress_threads NUM # Number of progress threads (default: 0) --enable_vmm # Enable VMM memory allocation when DRAM is requested ---large_blk_iter_ftr NUM # Factor to reduce transfer iteration for block size above 1MB (default: 16) +``` + +#### Device and Network Configuration +``` +--device_list LIST # Comma-separated device names (default: all) +--etcd_endpoints URL # ETCD server URL for coordination (default: http://localhost:2379) +``` + +#### Storage Backend Options (GDS, GDS_MT, POSIX, HF3FS, OBJ) +``` +--filepath PATH # File path for storage operations +--num_files NUM # Number of files used by benchmark (default: 1) +--storage_enable_direct # Enable direct I/O for storage operations +``` + +#### GDS Backend Specific Options +``` +--gds_batch_pool_size NUM # Batch pool size for GDS operations (default: 32) +--gds_batch_limit NUM # Batch limit for GDS operations (default: 128) +``` + +#### GDS_MT Backend Specific Options +``` +--gds_mt_num_threads NUM # Number of threads used by GDS MT plugin (default: 1) +``` + +#### POSIX Backend Specific Options +``` +--posix_api_type TYPE # API type for POSIX operations [AIO, URING] (default: AIO) +``` + +#### GPUNETIO Backend Specific Options +``` +--gpunetio_device_list LIST # Comma-separated GPU CUDA device id for GPUNETIO +``` + +#### OBJ (S3) Backend Specific Options +``` +--obj_access_key KEY # Access key for S3 backend +--obj_secret_key KEY # Secret key for S3 backend +--obj_session_token TOKEN # Session token for S3 backend +--obj_bucket_name NAME # Bucket name for S3 backend +--obj_scheme SCHEME # HTTP scheme for S3 backend [http, https] (default: http) +--obj_region REGION # Region for S3 backend (default: eu-central-1) +--obj_use_virtual_addressing # Use virtual addressing for S3 backend +--obj_endpoint_override URL # Endpoint override for S3 backend +--obj_req_checksum TYPE # Required checksum for S3 backend [supported, required] (default: supported) ``` ### Using ETCD for Coordination @@ -118,31 +221,97 @@ apt install etcd-server Example: ```bash # On host 1 -./nixlbench --runtime_type=ETCD --etcd-endpoints http://etcd-server:2379 --backend UCX --seg_type VRAM +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX --initiator_seg_type VRAM --target_seg_type VRAM # On host 2 -./nixlbench --runtime_type=ETCD --etcd-endpoints http://etcd-server:2379 --backend UCX --seg_type VRAM +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX --initiator_seg_type VRAM --target_seg_type VRAM ``` The workers automatically coordinate ranks through ETCD as they connect. -### Benchmarking the OBJ Plugin +### Backend-Specific Examples + +#### Network Backends + +**UCX Backend (Default)** +```bash +# Basic UCX benchmark +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX + +# UCX with specific devices +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX --device_list mlx5_0,mlx5_1 + +# UCX Memory-Only variant +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend UCX_MO +``` + +**GPUNETIO Backend** +```bash +# DOCA GPUNetIO with specific GPU devices +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend GPUNETIO --gpunetio_device_list 0,1 +``` + +#### Storage Backends + +**GDS (GPU Direct Storage)** +```bash +# Basic GDS benchmark +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend GDS --filepath /mnt/storage/testfile --storage_enable_direct + +# GDS with custom batch settings +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend GDS --filepath /mnt/storage/testfile --gds_batch_pool_size 64 --gds_batch_limit 256 +``` + +**GDS_MT (Multi-threaded GDS)** +```bash +# Multi-threaded GDS +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend GDS_MT --filepath /mnt/storage/testfile --gds_mt_num_threads 8 +``` + +**POSIX Backend** +```bash +# POSIX with AIO +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend POSIX --filepath /mnt/storage/testfile --posix_api_type AIO + +# POSIX with io_uring +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend POSIX --filepath /mnt/storage/testfile --posix_api_type URING --storage_enable_direct +``` + +#### Worker Types + +**NVSHMEM Worker** +```bash +# NVSHMEM (GPU-only, VRAM required) +./nixlbench --etcd_endpoints http://etcd-server:2379 --worker_type nvshmem --initiator_seg_type VRAM --target_seg_type VRAM +``` + +### Benchmarking the OBJ (S3) Backend For OBJ plugin benchmarking run etcd-server and a single nixlbench instance. Example: ```bash -AWS_ACCESS_KEY_ID= AWS_SECRET_ACCESS_KEY= AWS_DEFAULT_REGION= /tmp/nixlbench/nixlbench --etcd-endpoints http://:2379 --backend OBJ --obj_bucket_name +# Basic S3 benchmark using environment variables +AWS_ACCESS_KEY_ID= AWS_SECRET_ACCESS_KEY= AWS_DEFAULT_REGION= \ +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend OBJ --obj_bucket_name + +# S3 benchmark using command line flags +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend OBJ \ + --obj_access_key \ + --obj_secret_key \ + --obj_region \ + --obj_bucket_name ``` -Access key, secret access key, default region and bucket name are mandatory fields. -Use your own valid credentials. -Transfer times are higher than local storage so it is advisable to use less iterations than the default values. +**Performance Considerations:** +Transfer times are higher than local storage, so consider reducing iterations: -Example: ```bash -AWS_ACCESS_KEY_ID= AWS_SECRET_ACCESS_KEY= AWS_DEFAULT_REGION= /tmp/nixlbench/nixlbench --etcd-endpoints http://etcd-server:2379 --backend OBJ --obj_bucket_name nixl-ci-test --warmup_iter 32 --num_iter 32 --large_blk_iter_ftr 2 +./nixlbench --etcd_endpoints http://etcd-server:2379 --backend OBJ \ + --obj_bucket_name test-bucket \ + --warmup_iter 32 --num_iter 32 --large_blk_iter_ftr 2 ``` -The default benchmark command tests write. To test read ops add the flag: `--op_type READ`. -To test tranfer data validity add the flag: `--check_consistency true`. +**Testing Options:** +- Test read operations: `--op_type READ` +- Validate data consistency: `--check_consistency` diff --git a/benchmark/nixlbench/contrib/Dockerfile b/benchmark/nixlbench/contrib/Dockerfile index 13538ed96b..dff89c482c 100644 --- a/benchmark/nixlbench/contrib/Dockerfile +++ b/benchmark/nixlbench/contrib/Dockerfile @@ -50,22 +50,16 @@ RUN apt-get update -y && \ libgtest-dev \ build-essential -# Add Mellanox repository and install packages -RUN ARCH_SUFFIX=$(if [ "${ARCH}" = "aarch64" ]; then echo "arm64-sbsa"; else echo "${ARCH}"; fi) && \ - export PKG_CONFIG_PATH="/opt/mellanox/doca/lib/${ARCH_SUFFIX}-linux-gnu/pkgconfig:/opt/mellanox/dpdk/lib/${ARCH_SUFFIX}-linux-gnu/pkgconfig:$PKG_CONFIG_PATH" && \ - curl -fsSL https://linux.mellanox.com/public/repo/doca/3.0.0/ubuntu24.04/${ARCH_SUFFIX}/GPG-KEY-Mellanox.pub | \ - gpg --dearmor | tee /usr/share/keyrings/mellanox-archive-keyring.gpg && \ - echo "deb [signed-by=/usr/share/keyrings/mellanox-archive-keyring.gpg] https://linux.mellanox.com/public/repo/doca/3.0.0/ubuntu24.04/${ARCH_SUFFIX} ./" | \ - tee /etc/apt/sources.list.d/mellanox.list && \ - DEBIAN_FRONTEND=noninteractive apt update -y && \ - apt install -y --no-install-recommends \ - mlnx-dpdk mlnx-dpdk-dev \ - doca-sdk-common doca-sdk-dma doca-sdk-dpdk-bridge \ - doca-sdk-eth doca-sdk-flow doca-sdk-rdma doca-all \ - doca-sdk-gpunetio libdoca-sdk-gpunetio-dev +# Add DOCA repository and install packages +RUN ARCH_SUFFIX=$(if [ "${ARCH}" = "aarch64" ]; then echo "arm64"; else echo "amd64"; fi) && \ + MELLANOX_OS="$(. /etc/lsb-release; echo ${DISTRIB_ID}${DISTRIB_RELEASE} | tr A-Z a-z | tr -d .)" && \ + wget --tries=3 --waitretry=5 https://www.mellanox.com/downloads/DOCA/DOCA_v3.1.0/host/doca-host_3.1.0-091000-25.07-${MELLANOX_OS}_${ARCH_SUFFIX}.deb -O doca-host.deb && \ + dpkg -i doca-host.deb && \ + apt-get update && \ + apt-get install -y --no-install-recommends doca-sdk-gpunetio libdoca-sdk-gpunetio-dev libdoca-sdk-verbs-dev # Install AWS CLI -RUN curl "https://awscli.amazonaws.com/awscli-exe-linux-x86_64.zip" -o "awscliv2.zip" && \ +RUN curl "https://awscli.amazonaws.com/awscli-exe-linux-${ARCH}.zip" -o "awscliv2.zip" && \ unzip awscliv2.zip && ./aws/install && rm -rf awscliv2.zip aws # --- Stage 2a: Represents using UCX from the base image --- @@ -74,6 +68,8 @@ RUN echo "INFO: Using UCX from base image (UCX=${UCX})." # --- Stage 2b: Represents building UCX from source --- FROM os_setup_stage AS ucx_custom_image +ARG BUILD_TYPE="release" +ARG NPROC RUN mkdir -p /workspace/ucx COPY --from=ucx . /workspace/ucx @@ -89,9 +85,10 @@ RUN echo "INFO: Starting custom UCX build..." && \ ./autogen.sh && \ echo "INFO: Building UCX..." && \ ./contrib/configure-release --with-cuda=/usr/local/cuda \ + $(if [ "$BUILD_TYPE" = "debug" ]; then echo "--enable-debug"; fi) \ --enable-mt \ --without-go && \ - make -j$(nproc) && \ + make -j${NPROC:-$(nproc)} && \ make install && \ cd / && \ echo "INFO: Finished building and installing UCX." @@ -106,21 +103,32 @@ ARG ARCH="x86_64" ARG DEFAULT_PYTHON_VERSION ARG WHL_PYTHON_VERSIONS="3.12" ARG WHL_PLATFORM="manylinux_2_39_$ARCH" +ARG BUILD_TYPE="release" +ARG EFA_INSTALLER_VERSION="latest" +ARG EFA_INSTALL_PATH="/opt/amazon/efa" +ARG NPROC WORKDIR /workspace -RUN git clone https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ +RUN git clone --depth 1 https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ cd etcd-cpp-apiv3 && \ sed -i '/^find_dependency(cpprestsdk)$/d' etcd-cpp-api-config.in.cmake && \ mkdir build && cd build && \ - cmake .. -DBUILD_ETCD_CORE_ONLY=ON -DCMAKE_BUILD_TYPE=Release && make -j$(nproc) && make install + cmake .. -DBUILD_ETCD_CORE_ONLY=ON -DCMAKE_BUILD_TYPE=Release && make -j${NPROC:-$(nproc)} && make install COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/ COPY --from=nixl . /workspace/nixl COPY --from=nixlbench . /workspace/nixlbench # Install AWS SDK C++ dependencies and build -RUN apt-get update && apt-get install -y libcurl4-openssl-dev libssl-dev uuid-dev zlib1g-dev -RUN git clone --recurse-submodules https://github.com/aws/aws-sdk-cpp.git --branch 1.11.581 && \ +RUN apt-get update && apt-get install -y libcurl4-openssl-dev libssl-dev uuid-dev zlib1g-dev hwloc libhwloc-dev + +# Install EFA (Elastic Fabric Adapter) +RUN curl -fsSL "https://efa-installer.amazonaws.com/aws-efa-installer-${EFA_INSTALLER_VERSION}.tar.gz" | tar xz && \ + cd aws-efa-installer && \ + ./efa_installer.sh -y -g --skip-kmod --skip-limit-conf --no-verify && \ + ldconfig + +RUN git clone --recurse-submodules --depth 1 --shallow-submodules https://github.com/aws/aws-sdk-cpp.git --branch 1.11.581 && \ mkdir sdk_build && \ cd sdk_build && \ cmake ../aws-sdk-cpp/ -DCMAKE_BUILD_TYPE=Release -DBUILD_ONLY="s3" -DENABLE_TESTING=OFF -DCMAKE_INSTALL_PREFIX=/usr/local && \ @@ -129,16 +137,19 @@ RUN git clone --recurse-submodules https://github.com/aws/aws-sdk-cpp.git --bran WORKDIR /workspace/nixl -ENV LD_LIBRARY_PATH=/usr/local/lib:$LD_LIBRARY_PATH +ENV LD_LIBRARY_PATH=/usr/local/lib:$EFA_INSTALL_PATH/lib:$LD_LIBRARY_PATH ENV VIRTUAL_ENV=/workspace/nixl/.venv RUN uv venv $VIRTUAL_ENV --python $DEFAULT_PYTHON_VERSION && \ # pybind11 pip install needed for ubuntu 22.04 uv pip install --upgrade meson pybind11 patchelf pyYAML click tabulate +RUN CUDA_SHORT_VERSION=cu$(echo $CUDA_VERSION | cut -d. -f1,2 | tr -d .) && \ + uv pip install torch --index https://download.pytorch.org/whl/$CUDA_SHORT_VERSION + RUN rm -rf build && \ mkdir build && \ - uv run meson setup build/ --prefix=/usr/local/nixl && \ + uv run meson setup build -Dlibfabric_path=$EFA_INSTALL_PATH --prefix=/usr/local/nixl --buildtype=$BUILD_TYPE && \ cd build && \ ninja && \ ninja install @@ -171,7 +182,7 @@ RUN ls -ll /workspace/nixlbench RUN rm -rf build && \ mkdir build && \ - uv run meson setup build -Dnixl_path=/usr/local/nixl/ -Dprefix=/usr/local/nixlbench && \ + uv run meson setup build -Dnixl_path=/usr/local/nixl/ -Dprefix=/usr/local/nixlbench --buildtype=$BUILD_TYPE && \ cd build && ninja && ninja install WORKDIR /workspace/nixl @@ -181,7 +192,7 @@ RUN ls -ll benchmark/kvbench # Install dependencies for benchmarks RUN uv pip install -e benchmark/kvbench -ENV PATH=/usr/local/nixlbench/bin:$PATH +ENV PATH=/usr/local/nixlbench/bin:/usr/local/nixl/bin:$PATH ENV LD_LIBRARY_PATH=/usr/local/nixlbench/lib:$LD_LIBRARY_PATH ENV PYTHON_PATH=/usr/local/nixlbench/lib/python3/dist-packages/nixlbench/ diff --git a/benchmark/nixlbench/contrib/build.sh b/benchmark/nixlbench/contrib/build.sh index 1bce3d86b9..c7e2495c63 100755 --- a/benchmark/nixlbench/contrib/build.sh +++ b/benchmark/nixlbench/contrib/build.sh @@ -24,6 +24,7 @@ NIXL_BENCH_BUILD_CONTEXT_ARGS="--build-context nixlbench=$BUILD_CONTEXT/" DOCKER_FILE="${SOURCE_DIR}/Dockerfile" UCX_SRC="" UCX_BUILD_CONTEXT_ARGS="" +BUILD_TYPE="release" commit_id=$(git rev-parse --short HEAD) # Get latest TAG and add COMMIT_ID for dev @@ -40,7 +41,7 @@ ARCH=$(uname -m) WHL_BASE=manylinux_2_39 WHL_PLATFORM=${WHL_BASE}_${ARCH} WHL_PYTHON_VERSIONS="3.12" -OS="ubuntu24" +NPROC=${NPROC:-$(nproc)} get_options() { while :; do @@ -65,6 +66,14 @@ get_options() { missing_requirement $1 fi ;; + --build-type) + if [ "$2" ]; then + BUILD_TYPE=$2 + shift + else + missing_requirement $1 + fi + ;; --nixl) if [ "$2" ]; then NIXL_BUILD_CONTEXT_ARGS="--build-context nixl=$2" @@ -155,6 +164,7 @@ show_build_options() { echo "Building NIXLBench Image" echo "NIXL Source: ${NIXL_SRC}" echo "UCX Source: ${UCX_SRC} (optional)" + echo "Build Type: ${BUILD_TYPE}" echo "Image Tag: ${TAG}" echo "Build Context: ${BUILD_CONTEXT}" echo "Build Context Args: ${BUILD_CONTEXT_ARGS}" @@ -170,8 +180,8 @@ show_help() { echo " [--base-image-tag base image tag]" echo " [--nixlbench path/to/nixlbench/source/dir]" echo " [--ucx path/to/ucx/source/dir]" + echo " [--build-type [debug|release] to select build type]" echo " [--no-cache disable docker build cache]" - echo " [--os [ubuntu24|ubuntu22] to select Ubuntu version]" echo " [--python-versions python versions to build for, comma separated]" echo " [--tag tag for image]" echo " [--arch [x86_64|aarch64] to select target architecture]" @@ -193,6 +203,8 @@ BUILD_ARGS+=" --build-arg BASE_IMAGE=$BASE_IMAGE --build-arg BASE_IMAGE_TAG=$BAS BUILD_ARGS+=" --build-arg WHL_PYTHON_VERSIONS=$WHL_PYTHON_VERSIONS" BUILD_ARGS+=" --build-arg WHL_PLATFORM=$WHL_PLATFORM" BUILD_ARGS+=" --build-arg ARCH=$ARCH" +BUILD_ARGS+=" --build-arg BUILD_TYPE=$BUILD_TYPE" +BUILD_ARGS+=" --build-arg NPROC=$NPROC" show_build_options diff --git a/benchmark/nixlbench/meson.build b/benchmark/nixlbench/meson.build index 750de9e087..1910e7ad31 100644 --- a/benchmark/nixlbench/meson.build +++ b/benchmark/nixlbench/meson.build @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -project('nixlbench', 'CPP', version: '0.5.0', +project('nixlbench', 'CPP', version: '0.6.0', default_options: ['buildtype=release', 'werror=true', 'cpp_std=c++17', @@ -23,6 +23,7 @@ project('nixlbench', 'CPP', version: '0.5.0', # set up some global vars for compiler, platform, configuration, etc. cpp = meson.get_compiler('cpp') +fs = import('fs') # Allow overriding paths through environment variables # CUDA @@ -60,8 +61,28 @@ endif cuda_available = false if cuda_lib_path == '' cuda_dep = dependency('cuda', required : false, modules : [ 'cudart', 'cuda' ]) + if not cuda_dep.found() + # Meson is not detecting CUDA reliably on ARM, fallback to default + cuda_home = run_command('bash', '-c', 'echo $CUDA_HOME', check: true).stdout().strip() + if cuda_home == '' + cuda_home = '/usr/local/cuda' + endif + cuda_lib = cuda_home + '/lib64' + cuda_inc = cuda_home + '/include' + cuda_stub = cuda_lib + '/stubs' + if fs.exists(cuda_lib) and fs.exists(cuda_inc) + cuda_dep = declare_dependency( + link_args : ['-L' + cuda_lib, '-L' + cuda_stub, '-lcuda', '-lcudart'], + include_directories : include_directories(cuda_inc)) + if cuda_dep.found() + message('Found CUDA installation through fallback method:', cuda_home) + endif + endif + endif if cuda_dep.found() cuda_available = true + else + warning('CUDA not found. VRAM support will be disabled.') endif else message('cuda lib path ', cuda_lib_path) @@ -82,10 +103,6 @@ if cuda_available endif endif - -# UCX -ucx_dep = dependency('ucx') - # GFlags gflags_dep = dependency('gflags', required: true) @@ -206,7 +223,7 @@ if nvshmem_available ] if etcd_available - nvcc_args += ['-letcd-cpp-api', '-lcpprest'] + nvcc_args += ['-letcd-cpp-api-core', '-lcpprest'] etcd_rt = meson.current_build_dir() + '/src/runtime/etcd/libetcd_rt.a.p/etcd_rt.cpp.o' nvcc_cmd_files += etcd_rt endif diff --git a/benchmark/nixlbench/meson_options.txt b/benchmark/nixlbench/meson_options.txt index 1e9eda15bf..136cfda3c6 100644 --- a/benchmark/nixlbench/meson_options.txt +++ b/benchmark/nixlbench/meson_options.txt @@ -13,9 +13,9 @@ # See the License for the specific language governing permissions and # limitations under the License. -option('cudapath_inc', type: 'string', value: '/usr/local/cuda/include', description: 'Include path for CUDA') -option('cudapath_lib', type: 'string', value: '/usr/local/cuda/lib64/', description: 'Library path for CUDA') -option('cudapath_stub', type: 'string', value: '/usr/local/cuda/lib64/stubs', description: 'Extra Stub path for CUDA') +option('cudapath_inc', type: 'string', value: '', description: 'Include path for CUDA') +option('cudapath_lib', type: 'string', value: '', description: 'Library path for CUDA') +option('cudapath_stub', type: 'string', value: '', description: 'Extra Stub path for CUDA') option('etcd_inc_path', type: 'string', value: '', description: 'Path to ETCD C++ Client includes') option('etcd_lib_path', type: 'string', value: '', description: 'Path to ETCD C++ Client library') option('nixl_path', type: 'string', value: '/usr/local', description: 'Path to NiXL') diff --git a/benchmark/nixlbench/src/main.cpp b/benchmark/nixlbench/src/main.cpp index 288b95833f..57137347fd 100644 --- a/benchmark/nixlbench/src/main.cpp +++ b/benchmark/nixlbench/src/main.cpp @@ -116,7 +116,12 @@ static int processBatchSizes(xferBenchWorker &worker, num_threads); if (worker.isTarget()) { - worker.exchangeIOV(local_trans_lists); + if (xferBenchConfig::isStorageBackend()) { + std::cerr << "storage backend should be always an initiator" << std::endl; + return EXIT_FAILURE; + } + + worker.exchangeIOV(local_trans_lists, block_size); worker.poll(block_size); if (xferBenchConfig::check_consistency && xferBenchConfig::op_type == XFERBENCH_OP_WRITE) { @@ -125,14 +130,13 @@ static int processBatchSizes(xferBenchWorker &worker, if (IS_PAIRWISE_AND_SG()) { // TODO: This is here just to call throughput reduction // Separate reduction and print - xferBenchUtils::printStats(true, block_size, batch_size, 0); + xferBenchUtils::printStats(true, block_size, batch_size, xferBenchStats()); } } else if (worker.isInitiator()) { - std::vector> remote_trans_lists(worker.exchangeIOV(local_trans_lists)); + std::vector> remote_trans_lists( + worker.exchangeIOV(local_trans_lists, block_size)); - auto result = worker.transfer(block_size, - local_trans_lists, - remote_trans_lists); + auto result = worker.transfer(block_size, local_trans_lists, remote_trans_lists); if (std::holds_alternative(result)) { return 1; } @@ -148,8 +152,8 @@ static int processBatchSizes(xferBenchWorker &worker, } } - xferBenchUtils::printStats(false, block_size, batch_size, - std::get(result)); + xferBenchUtils::printStats( + false, block_size, batch_size, std::get(result)); } } diff --git a/benchmark/nixlbench/src/runtime/etcd/etcd_rt.cpp b/benchmark/nixlbench/src/runtime/etcd/etcd_rt.cpp index 31e8219fce..d1563a8b2c 100644 --- a/benchmark/nixlbench/src/runtime/etcd/etcd_rt.cpp +++ b/benchmark/nixlbench/src/runtime/etcd/etcd_rt.cpp @@ -412,6 +412,8 @@ int xferBenchEtcdRT::barrier(const std::string& barrier_id) { std::cerr << "Error in barrier " << e.what() << " " << barrier_key << " rank " << my_rank << " completed " << count << "/" << global_size << " ranks)" << std::endl; + // Clean up after failure, otherwise next runs may be affected + client->rmdir(barrier_key, true); } return -1; diff --git a/benchmark/nixlbench/src/runtime/etcd/test_etcd_runtime.py b/benchmark/nixlbench/src/runtime/etcd/test_etcd_runtime.py index d869dd99b4..20ffaaa920 100755 --- a/benchmark/nixlbench/src/runtime/etcd/test_etcd_runtime.py +++ b/benchmark/nixlbench/src/runtime/etcd/test_etcd_runtime.py @@ -22,15 +22,19 @@ import os import sys +from nixl.logging import get_logger + # Add the kvbench runtime path sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../../../kvbench/runtime")) +logger = get_logger(__name__) + try: from etcd_rt import _EtcdDistUtils def test_basic_functionality(): """Test basic rank and size functionality""" - print("Testing basic functionality...") + logger.info("Testing basic functionality...") # Initialize runtime - modify size based on how many processes you're running runtime = _EtcdDistUtils(etcd_endpoints="http://localhost:2379", size=2) @@ -38,37 +42,37 @@ def test_basic_functionality(): rank = runtime.get_rank() world_size = runtime.get_world_size() - print(f"Rank: {rank}, World Size: {world_size}") + logger.info("Rank: %d, World Size: %d", rank, world_size) # Test barrier - print(f"Rank {rank}: Before barrier") + logger.info("Rank %d: Before barrier", rank) runtime.barrier() - print(f"Rank {rank}: After barrier") + logger.info("Rank %d: After barrier", rank) # Test allgather my_data = {"rank": rank, "message": f"Hello from rank {rank}"} - print(f"Rank {rank}: Gathering data...") + logger.info("Rank %d: Gathering data...", rank) try: all_data = runtime.allgather_obj(my_data) - print(f"Rank {rank}: Gathered data from all ranks:") + logger.info("Rank %d: Gathered data from all ranks:", rank) for i, data in enumerate(all_data): - print(f" Rank {i}: {data}") + logger.info(" Rank %d: %s", i, data) except Exception as e: - print(f"Rank {rank}: Allgather failed: {e}") + logger.error("Rank %d: Allgather failed: %s", rank, e) # Test barrier again runtime.barrier() - print(f"Rank {rank}: Test completed successfully!") + logger.info("Rank %d: Test completed successfully!", rank) if __name__ == "__main__": test_basic_functionality() except ImportError as e: - print(f"Import error: {e}") - print("Make sure the etcd_runtime module is built and accessible") - print("Also ensure the etcd server is running at http://localhost:2379") + logger.error("Import error: %s", e) + logger.error("Make sure the etcd_runtime module is built and accessible") + logger.error("Also ensure the etcd server is running at http://localhost:2379") sys.exit(1) except Exception as e: - print(f"Runtime error: {e}") + logger.error("Runtime error: %s", e) sys.exit(1) diff --git a/benchmark/nixlbench/src/utils/utils.cpp b/benchmark/nixlbench/src/utils/utils.cpp index 489b2b12db..caa6b61056 100644 --- a/benchmark/nixlbench/src/utils/utils.cpp +++ b/benchmark/nixlbench/src/utils/utils.cpp @@ -15,8 +15,11 @@ * limitations under the License. */ +#include +#include #include #include +#include #include #include #include @@ -41,9 +44,10 @@ DEFINE_string(benchmark_group, "(Default: default)"); DEFINE_string(runtime_type, XFERBENCH_RT_ETCD, "Runtime type to use for communication [ETCD]"); DEFINE_string(worker_type, XFERBENCH_WORKER_NIXL, "Type of worker [nixl, nvshmem]"); -DEFINE_string(backend, - XFERBENCH_BACKEND_UCX, - "Name of communication backend [UCX, UCX_MO, GDS, POSIX, GPUNETIO, OBJ] \ +DEFINE_string( + backend, + XFERBENCH_BACKEND_UCX, + "Name of NIXL backend [UCX, UCX_MO, GDS, GDS_MT, POSIX, GPUNETIO, Mooncake, HF3FS, OBJ] \ (only used with nixl worker)"); DEFINE_string(initiator_seg_type, XFERBENCH_SEG_TYPE_DRAM, "Type of memory segment for initiator \ [DRAM, VRAM]"); @@ -75,9 +79,10 @@ DEFINE_int32 ( DEFINE_int32(num_initiator_dev, 1, "Number of device in initiator process"); DEFINE_int32(num_target_dev, 1, "Number of device in target process"); DEFINE_bool(enable_pt, false, "Enable Progress Thread (only used with nixl worker)"); +DEFINE_uint64(progress_threads, 0, "Number of progress threads (default: 0)"); DEFINE_bool(enable_vmm, false, "Enable VMM memory allocation when DRAM is requested"); -// Storage backend(GDS, POSIX, HF3FS, OBJ) options +// Storage backend(GDS, GDS_MT, POSIX, HF3FS, OBJ) options DEFINE_string (filepath, "", "File path for storage operations"); DEFINE_int32 (num_files, 1, "Number of files used by benchmark"); DEFINE_bool (storage_enable_direct, false, "Enable direct I/O for storage operations"); @@ -85,6 +90,7 @@ DEFINE_bool (storage_enable_direct, false, "Enable direct I/O for storage operat // GDS options - only used when backend is GDS DEFINE_int32(gds_batch_pool_size, 32, "Batch pool size for GDS operations (default: 32, only used with GDS backend)"); DEFINE_int32(gds_batch_limit, 128, "Batch limit for GDS operations (default: 128, only used with GDS backend)"); +DEFINE_int32(gds_mt_num_threads, 1, "Number of threads used by GDS MT plugin (Default: 1)"); // TODO: We should take rank wise device list as input to extend support // :, ... @@ -115,6 +121,9 @@ DEFINE_string(obj_req_checksum, XFERBENCH_OBJ_REQ_CHECKSUM_SUPPORTED, "Required checksum for S3 backend [supported, required]"); +// HF3FS options - only used when backend is HF3FS +DEFINE_int32(hf3fs_iopool_size, 64, "Size of io memory pool"); + std::string xferBenchConfig::runtime_type = ""; std::string xferBenchConfig::worker_type = ""; std::string xferBenchConfig::backend = ""; @@ -136,12 +145,14 @@ int xferBenchConfig::large_blk_iter_ftr = 16; int xferBenchConfig::warmup_iter = 0; int xferBenchConfig::num_threads = 0; bool xferBenchConfig::enable_pt = false; +size_t xferBenchConfig::progress_threads = 0; bool xferBenchConfig::enable_vmm = false; std::string xferBenchConfig::device_list = ""; std::string xferBenchConfig::etcd_endpoints = ""; std::string xferBenchConfig::benchmark_group = "default"; int xferBenchConfig::gds_batch_pool_size = 0; int xferBenchConfig::gds_batch_limit = 0; +int xferBenchConfig::gds_mt_num_threads = 0; std::string xferBenchConfig::gpunetio_device_list = ""; std::vector devices = { }; int xferBenchConfig::num_files = 0; @@ -158,6 +169,7 @@ std::string xferBenchConfig::obj_region = ""; bool xferBenchConfig::obj_use_virtual_addressing = false; std::string xferBenchConfig::obj_endpoint_override = ""; std::string xferBenchConfig::obj_req_checksum = ""; +int xferBenchConfig::hf3fs_iopool_size = 0; int xferBenchConfig::loadFromFlags() { @@ -169,6 +181,7 @@ xferBenchConfig::loadFromFlags() { if (worker_type == XFERBENCH_WORKER_NIXL) { backend = FLAGS_backend; enable_pt = FLAGS_enable_pt; + progress_threads = FLAGS_progress_threads; device_list = FLAGS_device_list; enable_vmm = FLAGS_enable_vmm; @@ -182,13 +195,15 @@ xferBenchConfig::loadFromFlags() { if (backend == XFERBENCH_BACKEND_GDS) { gds_batch_pool_size = FLAGS_gds_batch_pool_size; gds_batch_limit = FLAGS_gds_batch_limit; - storage_enable_direct = FLAGS_storage_enable_direct; + } + + if (backend == XFERBENCH_BACKEND_GDS_MT) { + gds_mt_num_threads = FLAGS_gds_mt_num_threads; } // Load POSIX-specific configurations if backend is POSIX if (backend == XFERBENCH_BACKEND_POSIX) { posix_api_type = FLAGS_posix_api_type; - storage_enable_direct = FLAGS_storage_enable_direct; // Validate POSIX API type if (posix_api_type != XFERBENCH_POSIX_API_AIO && @@ -206,7 +221,7 @@ xferBenchConfig::loadFromFlags() { // Load HD3FS-specific configurations if backend is HD3FS if (backend == XFERBENCH_BACKEND_HF3FS) { - storage_enable_direct = FLAGS_storage_enable_direct; + hf3fs_iopool_size = FLAGS_hf3fs_iopool_size; } // Load OBJ-specific configurations if backend is OBJ @@ -294,15 +309,16 @@ xferBenchConfig::loadFromFlags() { << std::endl; return -1; } - if (max_block_size > (total_buffer_size / num_threads)) { - std::cerr << "Incorrect buffer size configuration" << " max_block_size(" << max_block_size - << ") >" << " (total_buffer_size / num_threads)(" + if ((max_block_size * max_batch_size) > (total_buffer_size / num_threads)) { + std::cerr << "Incorrect buffer size configuration " << "(max_block_size * max_batch_size) " + << "(" << (max_block_size * max_batch_size) << ")" + << " is > (total_buffer_size / num_threads) (" << (total_buffer_size / num_threads) << ")" << std::endl; return -1; } - if (large_blk_iter_ftr == 0 || large_blk_iter_ftr > num_iter) { - std::cerr << "iter_factor must not be 0 and must be lower than num_iter" << std::endl; + if (large_blk_iter_ftr <= 0) { + std::cerr << "iter_factor must be greater than 0" << std::endl; return -1; } @@ -338,22 +354,30 @@ xferBenchConfig::loadFromFlags() { } void -xferBenchConfig::printOption (const std::string &desc, const std::string &value) { - std::cout << std::left << std::setw (60) << desc << ": " << value << std::endl; +xferBenchConfig::printOption(const std::string &desc, const std::string &value) { + std::cout << std::left << std::setw(60) << desc << ": " << value << std::endl; } -void xferBenchConfig::printConfig() { - std::cout << std::string(70, '*') << std::endl; +void +xferBenchConfig::printSeparator(const char sep) { + std::cout << std::string(160, sep) << std::endl; +} + +void +xferBenchConfig::printConfig() { + printSeparator('*'); std::cout << "NIXLBench Configuration" << std::endl; - std::cout << std::string(70, '*') << std::endl; - printOption ("Runtime (--runtime_type=[etcd])", runtime_type); + printSeparator('*'); + printOption("Runtime (--runtime_type=[etcd])", runtime_type); if (runtime_type == XFERBENCH_RT_ETCD) { - printOption ("ETCD Endpoint ", etcd_endpoints); + printOption("ETCD Endpoint ", etcd_endpoints); } - printOption ("Worker type (--worker_type=[nixl,nvshmem])", worker_type); + printOption("Worker type (--worker_type=[nixl,nvshmem])", worker_type); if (worker_type == XFERBENCH_WORKER_NIXL) { - printOption("Backend (--backend=[UCX,UCX_MO,GDS,POSIX,OBJ])", backend); + printOption("Backend (--backend=[UCX,UCX_MO,GDS,GDS_MT,POSIX,Mooncake,HF3FS,OBJ])", + backend); printOption ("Enable pt (--enable_pt=[0,1])", std::to_string (enable_pt)); + printOption("Progress threads (--progress_threads=N)", std::to_string(progress_threads)); printOption ("Device list (--device_list=dev1,dev2,...)", device_list); printOption ("Enable VMM (--enable_vmm=[0,1])", std::to_string (enable_vmm)); @@ -364,6 +388,11 @@ void xferBenchConfig::printConfig() { printOption ("GDS batch limit (--gds_batch_limit=N)", std::to_string (gds_batch_limit)); } + if (backend == XFERBENCH_BACKEND_GDS_MT) { + printOption("GDS MT Number of threads (--gds_mt_num_threads=N)", + std::to_string(gds_mt_num_threads)); + } + // Print POSIX options if backend is POSIX if (backend == XFERBENCH_BACKEND_POSIX) { printOption ("POSIX API type (--posix_api_type=[AIO,URING])", posix_api_type); @@ -417,7 +446,7 @@ void xferBenchConfig::printConfig() { printOption("Large block iter factor (--large_blk_iter_ftr=N)", std::to_string(large_blk_iter_ftr)); printOption ("Num threads (--num_threads=N)", std::to_string (num_threads)); - std::cout << std::string(80, '-') << std::endl; + printSeparator('-'); std::cout << std::endl; } @@ -450,6 +479,7 @@ std::vector xferBenchConfig::parseDeviceList() { bool xferBenchConfig::isStorageBackend() { return (XFERBENCH_BACKEND_GDS == xferBenchConfig::backend || + XFERBENCH_BACKEND_GDS_MT == xferBenchConfig::backend || XFERBENCH_BACKEND_HF3FS == xferBenchConfig::backend || XFERBENCH_BACKEND_POSIX == xferBenchConfig::backend || XFERBENCH_BACKEND_OBJ == xferBenchConfig::backend); @@ -587,32 +617,51 @@ void xferBenchUtils::checkConsistency(std::vector> &io } } -void xferBenchUtils::printStatsHeader() { +void +xferBenchUtils::printStatsHeader() { if (IS_PAIRWISE_AND_SG() && rt->getSize() > 2) { - std::cout << std::left << std::setw(20) << "Block Size (B)" + // clang-format off + std::cout << std::left + << std::setw(20) << "Block Size (B)" << std::setw(15) << "Batch Size" - << std::setw(15) << "Avg Lat. (us)" - << std::setw(15) << "B/W (MiB/Sec)" - << std::setw(15) << "B/W (GiB/Sec)" << std::setw(15) << "B/W (GB/Sec)" << std::setw(25) << "Aggregate B/W (GB/Sec)" << std::setw(20) << "Network Util (%)" + << std::setw(15) << "Avg Lat. (us)" + << std::setw(15) << "Avg Prep (us)" + << std::setw(15) << "P99 Prep (us)" + << std::setw(15) << "Avg Post (us)" + << std::setw(15) << "P99 Post (us)" + << std::setw(15) << "Avg Tx (us)" + << std::setw(15) << "P99 Tx (us)" << std::endl; + // clang-format on } else { - std::cout << std::left << std::setw(20) << "Block Size (B)" + // clang-format off + std::cout << std::left + << std::setw(20) << "Block Size (B)" << std::setw(15) << "Batch Size" - << std::setw(15) << "Avg Lat. (us)" - << std::setw(15) << "B/W (MiB/Sec)" - << std::setw(15) << "B/W (GiB/Sec)" << std::setw(15) << "B/W (GB/Sec)" + << std::setw(15) << "Avg Lat. (us)" + << std::setw(15) << "Avg Prep (us)" + << std::setw(15) << "P99 Prep (us)" + << std::setw(15) << "Avg Post (us)" + << std::setw(15) << "P99 Post (us)" + << std::setw(15) << "Avg Tx (us)" + << std::setw(15) << "P99 Tx (us)" << std::endl; + // clang-format on } - std::cout << std::string(80, '-') << std::endl; + xferBenchConfig::printSeparator('-'); } -void xferBenchUtils::printStats(bool is_target, size_t block_size, size_t batch_size, double total_duration) { +void +xferBenchUtils::printStats(bool is_target, + size_t block_size, + size_t batch_size, + xferBenchStats stats) { size_t total_data_transferred = 0; - double avg_latency = 0, throughput = 0, throughput_gib = 0, throughput_gb = 0; + double avg_latency = 0, throughput_gb = 0; double totalbw = 0; int num_iter = xferBenchConfig::num_iter; @@ -622,12 +671,15 @@ void xferBenchUtils::printStats(bool is_target, size_t block_size, size_t batch_ } // TODO: We can avoid this by creating a sub-communicator across initiator ranks - // if (isTarget() && IS_PAIRWISE_AND_SG() && rt->getSize() > 2) { - Fix this isTarget can not be called here + // if (isTarget() && IS_PAIRWISE_AND_SG() && rt->getSize() > 2) { - Fix this isTarget can not be + // called here if (is_target && IS_PAIRWISE_AND_SG() && rt->getSize() > 2) { rt->reduceSumDouble(&throughput_gb, &totalbw, 0); return; } + double total_duration = stats.total_duration.avg(); + total_data_transferred = ((block_size * batch_size) * num_iter); // In Bytes avg_latency = (total_duration / (num_iter * batch_size)); // In microsec if (IS_PAIRWISE_AND_MG()) { @@ -635,9 +687,6 @@ void xferBenchUtils::printStats(bool is_target, size_t block_size, size_t batch_ avg_latency /= xferBenchConfig::num_initiator_dev; // In microsec } - throughput = (((double) total_data_transferred / (1024 * 1024)) / - (total_duration / 1e6)); // In MiB/Sec - throughput_gib = (throughput / 1024); // In GiB/Sec throughput_gb = (((double) total_data_transferred / (1000 * 1000 * 1000)) / (total_duration / 1e6)); // In GB/Sec @@ -651,25 +700,48 @@ void xferBenchUtils::printStats(bool is_target, size_t block_size, size_t batch_ return; } + double prepare_duration = stats.prepare_duration.avg(); + double prepare_p99_duration = stats.prepare_duration.p99(); + double post_duration = stats.post_duration.avg(); + double post_p99_duration = stats.post_duration.p99(); + double transfer_duration = stats.transfer_duration.avg(); + double transfer_p99_duration = stats.transfer_duration.p99(); + // Tabulate print with fixed width for each string if (IS_PAIRWISE_AND_SG() && rt->getSize() > 2) { - std::cout << std::left << std::setw(20) << block_size + // clang-format off + std::cout << std::left << std::fixed << std::setprecision(6) + << std::setw(20) << block_size << std::setw(15) << batch_size - << std::setw(15) << avg_latency - << std::setw(15) << throughput - << std::setw(15) << throughput_gib << std::setw(15) << throughput_gb << std::setw(25) << totalbw - << std::setw(20) << (totalbw / (rt->getSize()/2 * MAXBW))*100 + << std::setw(20) << (totalbw / (rt->getSize() / 2 * MAXBW)) * 100 + << std::setprecision(1) + << std::setw(15) << avg_latency + << std::setw(15) << prepare_duration + << std::setw(15) << prepare_p99_duration + << std::setw(15) << post_duration + << std::setw(15) << post_p99_duration + << std::setw(15) << transfer_duration + << std::setw(15) << transfer_p99_duration << std::endl; + // clang-format on } else { - std::cout << std::left << std::setw(20) << block_size + // clang-format off + std::cout << std::left << std::fixed << std::setprecision(6) + << std::setw(20) << block_size << std::setw(15) << batch_size - << std::setw(15) << avg_latency - << std::setw(15) << throughput - << std::setw(15) << throughput_gib << std::setw(15) << throughput_gb + << std::setprecision(1) + << std::setw(15) << avg_latency + << std::setw(15) << prepare_duration + << std::setw(15) << prepare_p99_duration + << std::setw(15) << post_duration + << std::setw(15) << post_p99_duration + << std::setw(15) << transfer_duration + << std::setw(15) << transfer_p99_duration << std::endl; + // clang-format on } } @@ -799,3 +871,111 @@ xferBenchUtils::rmObjS3(const std::string &name) { } return true; } + +/* + * xferMetricStats + */ + +double +xferMetricStats::min() const { + if (samples.empty()) return 0; + return *std::min_element(samples.begin(), samples.end()); +} + +double +xferMetricStats::max() const { + if (samples.empty()) return 0; + return *std::max_element(samples.begin(), samples.end()); +} + +double +xferMetricStats::avg() const { + if (samples.empty()) return 0; + return std::accumulate(samples.begin(), samples.end(), 0.0) / samples.size(); +} + +double +xferMetricStats::p90() { + if (samples.empty()) return 0; + std::sort(samples.begin(), samples.end()); + size_t index = samples.size() * 0.9; + return samples[std::min(index, samples.size() - 1)]; +} + +double +xferMetricStats::p95() { + if (samples.empty()) return 0; + std::sort(samples.begin(), samples.end()); + size_t index = samples.size() * 0.95; + return samples[std::min(index, samples.size() - 1)]; +} + +double +xferMetricStats::p99() { + if (samples.empty()) return 0; + std::sort(samples.begin(), samples.end()); + size_t index = samples.size() * 0.99; + return samples[std::min(index, samples.size() - 1)]; +} + +void +xferMetricStats::add(double value) { + samples.push_back(value); +} + +void +xferMetricStats::add(const xferMetricStats &other) { + samples.insert(samples.end(), other.samples.begin(), other.samples.end()); +} + +void +xferMetricStats::reserve(size_t n) { + samples.reserve(n); +} + +void +xferMetricStats::clear() { + samples.clear(); +} + +/* + * xferBenchStats + */ + +void +xferBenchStats::clear() { + total_duration.clear(); + prepare_duration.clear(); + post_duration.clear(); + transfer_duration.clear(); +} + +void +xferBenchStats::add(const xferBenchStats &other) { + total_duration.add(other.total_duration); + prepare_duration.add(other.prepare_duration); + post_duration.add(other.post_duration); + transfer_duration.add(other.transfer_duration); +} + +void +xferBenchStats::reserve(size_t n) { + total_duration.reserve(n); + prepare_duration.reserve(n); + post_duration.reserve(n); + transfer_duration.reserve(n); +} + +/* + * xferBenchTimer + */ + +xferBenchTimer::xferBenchTimer() : start_(nixlTime::getUs()) {} + +nixlTime::us_t +xferBenchTimer::lap() { + nixlTime::us_t now = nixlTime::getUs(); + nixlTime::us_t duration = now - start_; + start_ = now; + return duration; +} diff --git a/benchmark/nixlbench/src/utils/utils.h b/benchmark/nixlbench/src/utils/utils.h index b6ff371da5..d7b27b5cc2 100644 --- a/benchmark/nixlbench/src/utils/utils.h +++ b/benchmark/nixlbench/src/utils/utils.h @@ -19,12 +19,14 @@ #define __UTILS_H #include "config.h" +#include #include #include #include #include #include #include +#include #include "runtime/runtime.h" #if HAVE_CUDA @@ -55,7 +57,6 @@ // TODO: This is true for CX-7, need support for other CX cards and NVLink #define MAXBW 50.0 // 400 Gbps or 50 GB/sec #define LARGE_BLOCK_SIZE (1LL * (1 << 20)) -#define MIN_WARMUP_ITERS 8 #define XFERBENCH_INITIATOR_BUFFER_ELEMENT 0xbb #define XFERBENCH_TARGET_BUFFER_ELEMENT 0xaa @@ -66,7 +67,9 @@ // Backend types #define XFERBENCH_BACKEND_UCX "UCX" #define XFERBENCH_BACKEND_UCX_MO "UCX_MO" +#define XFERBENCH_BACKEND_LIBFABRIC "LIBFABRIC" #define XFERBENCH_BACKEND_GDS "GDS" +#define XFERBENCH_BACKEND_GDS_MT "GDS_MT" #define XFERBENCH_BACKEND_POSIX "POSIX" #define XFERBENCH_BACKEND_GPUNETIO "GPUNETIO" #define XFERBENCH_BACKEND_MOONCAKE "Mooncake" @@ -140,6 +143,7 @@ class xferBenchConfig { static int warmup_iter; static int num_threads; static bool enable_pt; + static size_t progress_threads; static std::string device_list; static std::string etcd_endpoints; static std::string benchmark_group; @@ -150,6 +154,7 @@ class xferBenchConfig { static bool storage_enable_direct; static int gds_batch_pool_size; static int gds_batch_limit; + static int gds_mt_num_threads; static std::string gpunetio_device_list; static long page_size; static std::string obj_access_key; @@ -161,16 +166,79 @@ class xferBenchConfig { static bool obj_use_virtual_addressing; static std::string obj_endpoint_override; static std::string obj_req_checksum; + static int hf3fs_iopool_size; - static int loadFromFlags(); - static void printConfig(); + static int + loadFromFlags(); static void - printOption (const std::string &desc, const std::string &value); - static std::vector parseDeviceList(); + printConfig(); + static void + printOption(const std::string &desc, const std::string &value); + static void + printSeparator(const char sep = '-'); + static std::vector + parseDeviceList(); static bool isStorageBackend(); }; +// Timer class for measuring durations at high resolution +class xferBenchTimer { +public: + xferBenchTimer(); + + // Return the elapsed time in microseconds + nixlTime::us_t + lap(); + +private: + nixlTime::us_t start_; +}; + +// Stats class for measuring arbitrary numeric metrics with multiple samples +class xferMetricStats { +public: + double + min() const; + double + max() const; + double + avg() const; + double + p90(); + double + p95(); + double + p99(); + + void + add(double value); + void + add(const xferMetricStats &other); + void + reserve(size_t n); + void + clear(); + +private: + std::vector samples; +}; + +// Stats class for measuring benchmark metrics +struct xferBenchStats { + xferMetricStats total_duration; + xferMetricStats prepare_duration; + xferMetricStats post_duration; + xferMetricStats transfer_duration; + + void + clear(); + void + add(const xferBenchStats &other); + void + reserve(size_t n); +}; + // Generic IOV descriptor class independent of NIXL class xferBenchIOV { public: @@ -181,8 +249,12 @@ class xferBenchIOV { unsigned long long handle; std::string metaInfo; - xferBenchIOV(uintptr_t a, size_t l, int d) : - addr(a), len(l), devId(d), padded_size(len), handle(0) {} + xferBenchIOV(uintptr_t a, size_t l, int d) + : addr(a), + len(l), + devId(d), + padded_size(len), + handle(0) {} xferBenchIOV(uintptr_t a, size_t l, int d, size_t p, unsigned long long h) : addr(a), len(l), devId(d), padded_size(p), handle(h) {} @@ -213,10 +285,12 @@ class xferBenchUtils { static bool rmObjS3(const std::string &name); - static void checkConsistency(std::vector> &desc_lists); - static void printStatsHeader(); - static void printStats(bool is_target, size_t block_size, size_t batch_size, - double total_duration); + static void + checkConsistency(std::vector> &desc_lists); + static void + printStatsHeader(); + static void + printStats(bool is_target, size_t block_size, size_t batch_size, xferBenchStats stats); }; #endif // __UTILS_H diff --git a/benchmark/nixlbench/src/worker/nixl/nixl_worker.cpp b/benchmark/nixlbench/src/worker/nixl/nixl_worker.cpp index a5046cbda1..81facb2e50 100644 --- a/benchmark/nixlbench/src/worker/nixl/nixl_worker.cpp +++ b/benchmark/nixlbench/src/worker/nixl/nixl_worker.cpp @@ -31,16 +31,13 @@ #include #include #include +#include #include #include #define ROUND_UP(value, granularity) \ ((((value) + (granularity) - 1) / (granularity)) * (granularity)) -static uintptr_t gds_running_ptr = 0x0; -static std::vector> gds_remote_iovs; -static std::vector> storage_remote_iovs; - #define CHECK_NIXL_ERROR(result, message) \ do { \ if (0 != result) { \ @@ -81,13 +78,16 @@ xferBenchNixlWorker::xferBenchNixlWorker(int *argc, char ***argv, std::vector 1 ? + nixl_thread_sync_t::NIXL_THREAD_SYNC_RW : + nixl_thread_sync_t::NIXL_THREAD_SYNC_DEFAULT; char hostname[256]; nixl_mem_list_t mems; std::vector plugins; rank = rt->getRank(); - nixlAgentConfig dev_meta(enable_pt); + nixlAgentConfig dev_meta(enable_pt, false, 0, sync_mode); agent = new nixlAgent(name, dev_meta); @@ -95,12 +95,13 @@ xferBenchNixlWorker::xferBenchNixlWorker(int *argc, char ***argv, std::vector= 1) { @@ -129,9 +132,20 @@ xferBenchNixlWorker::xferBenchNixlWorker(int *argc, char ***argv, std::vector -createFileFds(std::string name) { - std::vector fds; +static std::vector +createFileFds(std::string name, int num_files) { + std::vector fds; int flags = O_RDWR | O_CREAT; - int num_files = xferBenchConfig::num_files; if (!xferBenchConfig::isStorageBackend()) { std::cerr << "Unknown storage backend: " << xferBenchConfig::backend << std::endl; @@ -421,26 +440,48 @@ createFileFds(std::string name) { for (int i = 0; i < num_files; i++) { std::string file_name = file_path + file_name_prefix + name + "_" + std::to_string(i); - std::cout << "Creating " - << " file: " << file_name << std::endl; + std::cout << "Creating file: " << file_name << std::endl; + + uint64_t file_size = 0; + if (XFERBENCH_OP_READ == xferBenchConfig::op_type) { + struct stat st; + if (::stat(file_name.c_str(), &st) == 0) { + std::cout << "File " << file_name << " exists, size: " << st.st_size << std::endl; + file_size = st.st_size; + } else { + std::cout << "File " << file_name << " does not exist, will be created." + << std::endl; + } + } + int fd = open(file_name.c_str(), flags, 0744); if (fd < 0) { std::cerr << "Failed to open file: " << file_name << " with error: " << strerror(errno) << std::endl; for (int j = 0; j < i; j++) { - close(fds[j]); + close(fds[j].fd); } return {}; } - fds.push_back(fd); + fds.emplace_back(xferFileState{fd, file_size, 0}); } return fds; } std::optional -xferBenchNixlWorker::initBasicDescFile(size_t buffer_size, int fd, int mem_dev_id) { - auto ret = - std::optional(std::in_place, (uintptr_t)gds_running_ptr, buffer_size, fd); +xferBenchNixlWorker::initBasicDescFile(size_t buffer_size, xferFileState &fstate, int mem_dev_id) { + int fd = fstate.fd; + uint64_t start_offset = fstate.offset; + uint64_t end_offset = fstate.offset + buffer_size; + auto ret = std::optional(std::in_place, fstate.offset, buffer_size, fd); + + fstate.offset = end_offset; + + // If in READ mode, only write if the region is not already present in the file + if (XFERBENCH_OP_READ == xferBenchConfig::op_type && end_offset <= fstate.file_size) { + return ret; + } + // Fill up with data void *buf; AllocationType type = AllocationType::MALLOC; @@ -456,21 +497,24 @@ xferBenchNixlWorker::initBasicDescFile(size_t buffer_size, int fd, int mem_dev_i // File is always initialized with XFERBENCH_TARGET_BUFFER_ELEMENT memset(buf, XFERBENCH_TARGET_BUFFER_ELEMENT, buffer_size); - if (xferBenchConfig::storage_enable_direct) { - gds_running_ptr = - ((gds_running_ptr + xferBenchConfig::page_size - 1) / xferBenchConfig::page_size) * - xferBenchConfig::page_size; - } else { - gds_running_ptr += (buffer_size * mem_dev_id); - } - int rc = pwrite(fd, buf, buffer_size, gds_running_ptr); - if (rc < 0) { - std::cerr << "Failed to write to file: " << fd << " with error: " << strerror(errno) - << std::endl; - return std::nullopt; + + size_t offset = start_offset; + while (buffer_size > 0) { + ssize_t rc = pwrite(fd, buf, buffer_size, offset); + if (rc < 0) { + std::cerr << "Failed to write to file: " << fd << " with error: " << strerror(errno) + << std::endl; + return std::nullopt; + } + + buffer_size -= rc; + offset += rc; } + free(buf); + if (end_offset > fstate.file_size) fstate.file_size = end_offset; + return ret; } @@ -514,7 +558,7 @@ xferBenchNixlWorker::cleanupBasicDescObj(xferBenchIOV &iov) { } std::vector> -xferBenchNixlWorker::allocateMemory(int num_lists) { +xferBenchNixlWorker::allocateMemory(int num_threads) { std::vector> iov_lists; size_t i, buffer_size, num_devices = 0; nixl_opt_args_t opt_args; @@ -524,7 +568,7 @@ xferBenchNixlWorker::allocateMemory(int num_lists) { } else if (isTarget()) { num_devices = xferBenchConfig::num_target_dev; } - buffer_size = xferBenchConfig::total_buffer_size / (num_devices * num_lists); + buffer_size = xferBenchConfig::total_buffer_size / (num_devices * num_threads); if (xferBenchConfig::storage_enable_direct) { if (xferBenchConfig::page_size == 0) { @@ -543,7 +587,7 @@ xferBenchNixlWorker::allocateMemory(int num_lists) { gettimeofday(&tv, nullptr); uint64_t timestamp = tv.tv_sec * 1000000ULL + tv.tv_usec; - for (int list_idx = 0; list_idx < num_lists; list_idx++) { + for (int list_idx = 0; list_idx < num_threads; list_idx++) { std::vector iov_list; for (i = 0; i < num_devices; i++) { std::optional basic_desc; @@ -569,31 +613,50 @@ xferBenchNixlWorker::allocateMemory(int num_lists) { remote_iovs.push_back(iov_list); } } else if (xferBenchConfig::isStorageBackend()) { + int num_buffers = num_threads * num_devices; + int num_files = xferBenchConfig::num_files; + int remainder_buffers = num_buffers % num_files; + + if (num_files > num_buffers) { + std::cerr << "Error: number of buffers (" << num_buffers + << ") needs to be bigger or equal to the number of files (" << num_files + << "). Try adjusting num_files." << std::endl; + exit(EXIT_FAILURE); + } + + if (remainder_buffers != 0) { + std::cerr << "Error: number of buffers (" << num_buffers + << ") needs to be divisible by the number of files (" << num_files + << "). Try adjusting num_files." << std::endl; + exit(EXIT_FAILURE); + } - remote_fds = createFileFds(getName()); + remote_fds = createFileFds(getName(), num_files); if (remote_fds.empty()) { std::cerr << "Failed to create " << xferBenchConfig::backend << " file" << std::endl; exit(EXIT_FAILURE); } - for (int list_idx = 0; list_idx < num_lists; list_idx++) { + + int file_idx = 0; + for (int list_idx = 0; list_idx < num_threads; list_idx++) { std::vector iov_list; for (i = 0; i < num_devices; i++) { std::optional basic_desc; - basic_desc = initBasicDescFile(buffer_size, remote_fds[0], i); + basic_desc = initBasicDescFile(buffer_size, remote_fds[file_idx], i); if (basic_desc) { iov_list.push_back(basic_desc.value()); } + file_idx += 1; + if (file_idx >= num_files) file_idx = 0; } nixl_reg_dlist_t desc_list(FILE_SEG); iovListToNixlRegDlist(iov_list, desc_list); CHECK_NIXL_ERROR(agent->registerMem(desc_list, &opt_args), "registerMem failed"); remote_iovs.push_back(iov_list); } - // Reset the running pointer to 0 - gds_running_ptr = 0x0; } - for (int list_idx = 0; list_idx < num_lists; list_idx++) { + for (int list_idx = 0; list_idx < num_threads; list_idx++) { std::vector iov_list; for (i = 0; i < num_devices; i++) { std::optional basic_desc; @@ -613,7 +676,9 @@ xferBenchNixlWorker::allocateMemory(int num_lists) { } if (basic_desc) { - basic_desc.value().metaInfo = remote_iovs[list_idx][i].metaInfo; + if (!remote_iovs.empty()) { + basic_desc.value().metaInfo = remote_iovs[list_idx][i].metaInfo; + } iov_list.push_back(basic_desc.value()); } } @@ -703,7 +768,6 @@ xferBenchNixlWorker::exchangeMetadata() { rt->sendInt(&meta_sz, destrank); rt->sendChar((char *)buffer, meta_sz, destrank); } else if (isInitiator()) { - char *buffer; std::string remote_agent; int srcrank; @@ -714,37 +778,60 @@ xferBenchNixlWorker::exchangeMetadata() { } else { srcrank = 1; } - rt->recvInt(&meta_sz, srcrank); - buffer = (char *)calloc(meta_sz, sizeof(*buffer)); - rt->recvChar((char *)buffer, meta_sz, srcrank); - std::string remote_metadata(buffer, meta_sz); - agent->loadRemoteMD(remote_metadata, remote_agent); - if ("" == remote_agent) { - std::cerr << "NIXL: loadMetadata failed" << std::endl; + ret = rt->recvInt(&meta_sz, srcrank); + if (ret < 0) { + std::cerr << "NIXL: failed to receive metadata size" << std::endl; + return ret; + } + + std::string remote_metadata(meta_sz, '\0'); + ret = rt->recvChar(remote_metadata.data(), meta_sz, srcrank); + if (ret < 0) { + std::cerr << "NIXL: failed to receive metadata" << std::endl; + return ret; + } + + nixl_status_t status = agent->loadRemoteMD(remote_metadata, remote_agent); + if (status != NIXL_SUCCESS) { + std::cerr << "NIXL: loadRemoteMD failed: " << nixlEnumStrings::statusStr(status) + << std::endl; + return -1; } - free(buffer); } + return ret; } std::vector> -xferBenchNixlWorker::exchangeIOV(const std::vector> &local_iovs) { +xferBenchNixlWorker::exchangeIOV(const std::vector> &local_iovs, + size_t block_size) { std::vector> res; int desc_str_sz; if (xferBenchConfig::isStorageBackend()) { + size_t fd_idx = 0; + uint64_t file_offset = 0; for (auto &iov_list : local_iovs) { std::vector remote_iov_list; for (auto &iov : iov_list) { - std::optional basic_desc; if (XFERBENCH_BACKEND_OBJ == xferBenchConfig::backend) { + std::optional basic_desc; basic_desc = initBasicDescObj(iov.len, iov.devId, iov.metaInfo); + if (basic_desc) { + remote_iov_list.push_back(basic_desc.value()); + } } else { - basic_desc = initBasicDescFile(iov.len, remote_fds[0], iov.devId); - } - if (basic_desc) { - remote_iov_list.push_back(basic_desc.value()); + xferBenchIOV iov_remote(iov); + iov_remote.addr = file_offset; + iov_remote.len = block_size; + iov_remote.devId = remote_fds[fd_idx].fd; + remote_iov_list.push_back(iov_remote); + fd_idx++; + if (fd_idx >= remote_fds.size()) { + file_offset += block_size; + fd_idx = 0; + } } } res.push_back(remote_iov_list); @@ -757,14 +844,7 @@ xferBenchNixlWorker::exchangeIOV(const std::vector> &l iovListToNixlXferDlist(local_iov, local_desc); if (isTarget()) { - const char *buffer; int destrank; - - local_desc.serialize(&ser_des); - std::string desc_str = ser_des.exportStr(); - buffer = desc_str.data(); - desc_str_sz = desc_str.size(); - if (IS_PAIRWISE_AND_SG()) { destrank = rt->getRank() - xferBenchConfig::num_target_dev; // XXX: Fix up the rank, depends on processes distributed on hosts @@ -772,12 +852,14 @@ xferBenchNixlWorker::exchangeIOV(const std::vector> &l } else { destrank = 0; } + + local_desc.serialize(&ser_des); + std::string desc_str = ser_des.exportStr(); + desc_str_sz = desc_str.size(); rt->sendInt(&desc_str_sz, destrank); - rt->sendChar((char *)buffer, desc_str_sz, destrank); + rt->sendChar(desc_str.data(), desc_str.size(), destrank); } else if (isInitiator()) { - char *buffer; int srcrank; - if (IS_PAIRWISE_AND_SG()) { srcrank = rt->getRank() + xferBenchConfig::num_initiator_dev; // XXX: Fix up the rank, depends on processes distributed on hosts @@ -785,11 +867,19 @@ xferBenchNixlWorker::exchangeIOV(const std::vector> &l } else { srcrank = 1; } - rt->recvInt(&desc_str_sz, srcrank); - buffer = (char *)calloc(desc_str_sz, sizeof(*buffer)); - rt->recvChar((char *)buffer, desc_str_sz, srcrank); - std::string desc_str(buffer, desc_str_sz); + if (rt->recvInt(&desc_str_sz, srcrank) != 0) { + std::cerr << "NIXL: failed to receive metadata size" << std::endl; + std::exit(EXIT_FAILURE); + } + + std::string desc_str; + desc_str.resize(desc_str_sz, '\0'); + if (rt->recvChar(desc_str.data(), desc_str.size(), srcrank) != 0) { + std::cerr << "NIXL: failed to receive metadata" << std::endl; + std::exit(EXIT_FAILURE); + } + ser_des.importStr(desc_str); nixl_xfer_dlist_t remote_desc(&ser_des); @@ -808,11 +898,17 @@ execTransfer(nixlAgent *agent, const std::vector> &remote_iovs, const nixl_xfer_op_t op, const int num_iter, - const int num_threads) { + const int num_threads, + xferBenchStats &stats) { int ret = 0; + stats.clear(); + xferBenchTimer total_timer; #pragma omp parallel num_threads(num_threads) { + xferBenchStats thread_stats; + thread_stats.reserve(num_iter); + xferBenchTimer timer; const int tid = omp_get_thread_num(); const auto &local_iov = local_iovs[tid]; const auto &remote_iov = remote_iovs[tid]; @@ -832,16 +928,12 @@ execTransfer(nixlAgent *agent, nixl_opt_args_t params; nixl_b_params_t b_params; - bool error = false; nixlXferReqH *req; nixl_status_t rc; std::string target; if (xferBenchConfig::isStorageBackend()) { target = "initiator"; - } else if (XFERBENCH_BACKEND_MOONCAKE == xferBenchConfig::backend) { - params.hasNotif = false; - target = "target"; } else { params.notifMsg = "0xBEEF"; params.hasNotif = true; @@ -849,44 +941,52 @@ execTransfer(nixlAgent *agent, } CHECK_NIXL_ERROR(agent->createXferReq(op, local_desc, remote_desc, target, req, ¶ms), - "createTransferReq failed"); + "createXferReq failed"); + + const nixlTime::us_t prepare_duration = timer.lap(); + thread_stats.prepare_duration.add(prepare_duration); - for (int i = 0; i < num_iter && !error; i++) { + for (int i = 0; i < num_iter; i++) { rc = agent->postXferReq(req); - if (NIXL_ERR_BACKEND == rc) { - std::cout << "NIXL postRequest failed" << std::endl; - error = true; - } else { - do { - /* XXX agent isn't const because the getXferStatus() is not const */ - rc = agent->getXferStatus(req); - if (NIXL_ERR_BACKEND == rc) { - std::cout << "NIXL getStatus failed" << std::endl; - error = true; - break; - } - } while (NIXL_SUCCESS != rc); + const nixlTime::us_t post_duration = timer.lap(); + thread_stats.post_duration.add(post_duration); + while (NIXL_IN_PROG == rc) { + /* XXX agent isn't const because the getXferStatus() is not const */ + rc = agent->getXferStatus(req); + } + + if (NIXL_SUCCESS != rc) { + std::cout << "NIXL Xfer failed with status: " << nixlEnumStrings::statusStr(rc) + << std::endl; + ret = -1; + break; } + + const nixlTime::us_t transfer_duration = timer.lap(); + thread_stats.transfer_duration.add(transfer_duration); } - agent->releaseXferReq(req); - if (error) { + rc = agent->releaseXferReq(req); + if (NIXL_SUCCESS != rc) { std::cout << "NIXL releaseXferReq failed" << std::endl; ret = -1; } - } +#pragma omp critical + { stats.add(thread_stats); } + } + const nixlTime::us_t total_duration = total_timer.lap(); + stats.total_duration.add(total_duration); return ret; } -std::variant +std::variant xferBenchNixlWorker::transfer(size_t block_size, const std::vector> &local_iovs, const std::vector> &remote_iovs) { int num_iter = xferBenchConfig::num_iter / xferBenchConfig::num_threads; int skip = xferBenchConfig::warmup_iter / xferBenchConfig::num_threads; - struct timeval t_start, t_end; - double total_duration = 0.0; + xferBenchStats stats; int ret = 0; nixl_xfer_op_t xfer_op = XFERBENCH_OP_READ == xferBenchConfig::op_type ? NIXL_READ : NIXL_WRITE; // int completion_flag = 1; @@ -894,31 +994,30 @@ xferBenchNixlWorker::transfer(size_t block_size, // Reduce skip by 10x for large block sizes if (block_size > LARGE_BLOCK_SIZE) { skip /= xferBenchConfig::large_blk_iter_ftr; - if (skip < MIN_WARMUP_ITERS) { - skip = MIN_WARMUP_ITERS; - } num_iter /= xferBenchConfig::large_blk_iter_ftr; } - ret = execTransfer(agent, local_iovs, remote_iovs, xfer_op, skip, xferBenchConfig::num_threads); - if (ret < 0) { - return std::variant(ret); + if (skip > 0) { + ret = execTransfer( + agent, local_iovs, remote_iovs, xfer_op, skip, xferBenchConfig::num_threads, stats); + if (ret < 0) { + return std::variant(ret); + } } // Synchronize to ensure all processes have completed the warmup (iter and polling) synchronize(); - gettimeofday(&t_start, nullptr); + stats.clear(); ret = execTransfer( - agent, local_iovs, remote_iovs, xfer_op, num_iter, xferBenchConfig::num_threads); - - gettimeofday(&t_end, nullptr); - total_duration += - (((t_end.tv_sec - t_start.tv_sec) * 1e6) + (t_end.tv_usec - t_start.tv_usec)); // In us + agent, local_iovs, remote_iovs, xfer_op, num_iter, xferBenchConfig::num_threads, stats); + if (ret < 0) { + return std::variant(ret); + } synchronize(); - return ret < 0 ? std::variant(ret) : std::variant(total_duration); + return std::variant(stats); } void @@ -932,9 +1031,6 @@ xferBenchNixlWorker::poll(size_t block_size) { // Reduce skip by 10x for large block sizes if (block_size > LARGE_BLOCK_SIZE) { skip /= xferBenchConfig::large_blk_iter_ftr; - if (skip < MIN_WARMUP_ITERS) { - skip = MIN_WARMUP_ITERS; - } num_iter /= xferBenchConfig::large_blk_iter_ftr; } total_iter = skip + num_iter; diff --git a/benchmark/nixlbench/src/worker/nixl/nixl_worker.h b/benchmark/nixlbench/src/worker/nixl/nixl_worker.h index c2e9485f64..7cb8f62ef0 100644 --- a/benchmark/nixlbench/src/worker/nixl/nixl_worker.h +++ b/benchmark/nixlbench/src/worker/nixl/nixl_worker.h @@ -29,12 +29,18 @@ #include "utils/utils.h" #include "worker/worker.h" +struct xferFileState { + int fd; + uint64_t file_size; + uint64_t offset; +}; + class xferBenchNixlWorker: public xferBenchWorker { private: nixlAgent* agent; nixlBackendH* backend_engine; nixl_mem_t seg_type; - std::vector remote_fds; + std::vector remote_fds; std::vector> remote_iovs; public: xferBenchNixlWorker(int *argc, char ***argv, std::vector devices); @@ -46,24 +52,30 @@ class xferBenchNixlWorker: public xferBenchWorker { // Communication and synchronization int exchangeMetadata() override; - std::vector> exchangeIOV(const std::vector> - &local_iov_lists) override; - void poll(size_t block_size) override; - int synchronizeStart(); + std::vector> + exchangeIOV(const std::vector> &local_iov_lists, + size_t block_size) override; + void + poll(size_t block_size) override; + int + synchronizeStart(); // Data operations - std::variant transfer(size_t block_size, - const std::vector> &local_iov_lists, - const std::vector> &remote_iov_lists) override; + std::variant + transfer(size_t block_size, + const std::vector> &local_iov_lists, + const std::vector> &remote_iov_lists) override; private: - std::optional initBasicDescDram(size_t buffer_size, int mem_dev_id); + std::optional + initBasicDescDram(size_t buffer_size, int mem_dev_id); void cleanupBasicDescDram(xferBenchIOV &basic_desc); #if HAVE_CUDA std::optional initBasicDescVram(size_t buffer_size, int mem_dev_id); void cleanupBasicDescVram(xferBenchIOV &basic_desc); #endif - std::optional initBasicDescFile(size_t buffer_size, int fd, int mem_dev_id); + std::optional + initBasicDescFile(size_t buffer_size, xferFileState &fstate, int mem_dev_id); void cleanupBasicDescFile(xferBenchIOV &basic_desc); std::optional initBasicDescObj(size_t buffer_size, int mem_dev_id, std::string name); diff --git a/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.cpp b/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.cpp index bfe686d1fc..4c5efd1de9 100644 --- a/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.cpp +++ b/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.cpp @@ -116,16 +116,21 @@ int xferBenchNvshmemWorker::exchangeMetadata() { return 0; } -std::vector> xferBenchNvshmemWorker::exchangeIOV(const std::vector> &iov_lists) { +std::vector> +xferBenchNvshmemWorker::exchangeIOV(const std::vector> &iov_lists, + size_t block_size) { // For NVSHMEM, we don't need to exchange IOV lists // This will just return local IOV list return iov_lists; } // No thread support for NVSHMEM yet -static int execTransfer(const std::vector> &local_iovs, - const std::vector> &remote_iovs, - const int num_iter, cudaStream_t stream) { +static int +execTransfer(const std::vector> &local_iovs, + const std::vector> &remote_iovs, + const int num_iter, + cudaStream_t stream, + xferBenchStats &stats) { int ret = 0, tid = 0, target_rank; target_rank = 1; @@ -133,30 +138,41 @@ static int execTransfer(const std::vector> &local_iovs const auto &local_iov = local_iovs[tid]; const auto &remote_iov = remote_iovs[tid]; + xferBenchTimer total_timer; + xferBenchTimer timer; + for (int i = 0; i < num_iter; i++) { for (size_t i = 0; i < local_iov.size(); i++) { auto &local = local_iov[i]; auto &remote = remote_iov[i]; if (XFERBENCH_OP_WRITE == xferBenchConfig::op_type) { - nvshmemx_putmem_on_stream((void *)remote.addr, (void *)local.addr, local.len, target_rank, stream); + nvshmemx_putmem_on_stream( + (void *)remote.addr, (void *)local.addr, local.len, target_rank, stream); } else if (XFERBENCH_OP_READ == xferBenchConfig::op_type) { - nvshmemx_getmem_on_stream((void *)remote.addr, (void *)local.addr, local.len, target_rank, stream); + nvshmemx_getmem_on_stream( + (void *)remote.addr, (void *)local.addr, local.len, target_rank, stream); } } nvshmemx_quiet_on_stream(stream); + nixlTime::us_t transfer_duration = timer.lap(); + stats.transfer_duration.add(transfer_duration); } + nixlTime::us_t total_duration = total_timer.lap(); + stats.total_duration.add(total_duration); + return ret; } -std::variant xferBenchNvshmemWorker::transfer(size_t block_size, - const std::vector> &local_trans_lists, - const std::vector> &remote_trans_lists) { +std::variant +xferBenchNvshmemWorker::transfer(size_t block_size, + const std::vector> &local_trans_lists, + const std::vector> &remote_trans_lists) { cudaEvent_t start_event, stop_event; - float total_duration = 0.0; int num_iter = xferBenchConfig::num_iter / xferBenchConfig::num_threads; int skip = xferBenchConfig::warmup_iter / xferBenchConfig::num_threads; int ret = 0; + xferBenchStats stats; // Create events to time the transfer CHECK_CUDA_ERROR(cudaEventCreate(&start_event), "Failed to create CUDA event"); @@ -166,22 +182,22 @@ std::variant xferBenchNvshmemWorker::transfer(size_t block_size, // Reduce skip by 10x for large block sizes if (block_size > LARGE_BLOCK_SIZE) { skip /= xferBenchConfig::large_blk_iter_ftr; - if (skip < MIN_WARMUP_ITERS) { - skip = MIN_WARMUP_ITERS; - } num_iter /= xferBenchConfig::large_blk_iter_ftr; } - ret = execTransfer(local_trans_lists, remote_trans_lists, skip, stream); - if (ret < 0) { - return std::variant(ret); + if (skip > 0) { + ret = execTransfer(local_trans_lists, remote_trans_lists, skip, stream, stats); + if (ret < 0) { + return std::variant(ret); + } + stats.clear(); } nvshmemx_barrier_all_on_stream(stream); CHECK_CUDA_ERROR(cudaStreamSynchronize(stream), "Failed to synchronize CUDA stream"); CHECK_CUDA_ERROR(cudaEventRecord(start_event, stream), "Failed to record CUDA event"); - ret = execTransfer(local_trans_lists, remote_trans_lists, num_iter, stream); + ret = execTransfer(local_trans_lists, remote_trans_lists, num_iter, stream, stats); CHECK_CUDA_ERROR(cudaEventRecord(stop_event, stream), "Failed to record CUDA event"); @@ -189,13 +205,12 @@ std::variant xferBenchNvshmemWorker::transfer(size_t block_size, CHECK_CUDA_ERROR(cudaEventSynchronize(stop_event), "Failed to synchronize CUDA event"); CHECK_CUDA_ERROR(cudaStreamSynchronize(stream), "Failed to synchronize CUDA stream"); - // Time in ms - CHECK_CUDA_ERROR(cudaEventElapsedTime(&total_duration, start_event, stop_event), "Failed to get elapsed time"); - - return ret < 0 ? std::variant(ret) : std::variant((double)total_duration * 1e+3); + return ret < 0 ? std::variant(ret) : + std::variant(stats); } -void xferBenchNvshmemWorker::poll(size_t block_size) { +void +xferBenchNvshmemWorker::poll(size_t block_size) { // For NVSHMEM, we don't need to poll // The transfer is already complete when we reach this point nvshmemx_barrier_all_on_stream(stream); diff --git a/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.h b/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.h index c22c8f6aaf..95bffc9d9d 100644 --- a/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.h +++ b/benchmark/nixlbench/src/worker/nvshmem/nvshmem_worker.h @@ -47,18 +47,25 @@ class xferBenchNvshmemWorker: public xferBenchWorker { // Communication and synchronization int exchangeMetadata() override; - std::vector> exchangeIOV(const std::vector> - &local_iov_lists) override; - void poll(size_t block_size) override; - int synchronizeStart(); + std::vector> + exchangeIOV(const std::vector> &local_iov_lists, + size_t block_size) override; + void + poll(size_t block_size) override; + int + synchronizeStart(); // Data operations - std::variant transfer(size_t block_size, - const std::vector> &local_iov_lists, - const std::vector> &remote_iov_lists) override; + std::variant + transfer(size_t block_size, + const std::vector> &local_iov_lists, + const std::vector> &remote_iov_lists) override; + private: - std::optional initBasicDescNvshmem(size_t buffer_size, int mem_dev_id); - void cleanupBasicDescNvshmem(xferBenchIOV &iov); + std::optional + initBasicDescNvshmem(size_t buffer_size, int mem_dev_id); + void + cleanupBasicDescNvshmem(xferBenchIOV &iov); }; #endif diff --git a/benchmark/nixlbench/src/worker/worker.cpp b/benchmark/nixlbench/src/worker/worker.cpp index d51633a163..42e4c212f5 100644 --- a/benchmark/nixlbench/src/worker/worker.cpp +++ b/benchmark/nixlbench/src/worker/worker.cpp @@ -46,7 +46,14 @@ static xferBenchRT *createRT(int *terminate) { } int xferBenchWorker::synchronize() { - return rt->barrier("sync"); + if (rt->barrier("sync") != 0) { + std::cerr << "Failed to synchronize" << std::endl; + // assuming this is a fatal error, continue benchmarking after synchronization failure does + // not make sense + exit(EXIT_FAILURE); + } + + return 0; } xferBenchWorker::xferBenchWorker(int *argc, char ***argv) { diff --git a/benchmark/nixlbench/src/worker/worker.h b/benchmark/nixlbench/src/worker/worker.h index 6f090ba0cf..14d75af11d 100644 --- a/benchmark/nixlbench/src/worker/worker.h +++ b/benchmark/nixlbench/src/worker/worker.h @@ -49,15 +49,19 @@ class xferBenchWorker { // Communication and synchronization virtual int exchangeMetadata() = 0; - virtual std::vector> exchangeIOV(const std::vector> - &local_iov_lists) = 0; - virtual void poll(size_t block_size) = 0; - virtual int synchronizeStart() = 0; + virtual std::vector> + exchangeIOV(const std::vector> &local_iov_lists, + size_t block_size) = 0; + virtual void + poll(size_t block_size) = 0; + virtual int + synchronizeStart() = 0; // Data operations - virtual std::variant transfer(size_t block_size, - const std::vector> &local_iov_lists, - const std::vector> &remote_iov_lists) = 0; + virtual std::variant + transfer(size_t block_size, + const std::vector> &local_iov_lists, + const std::vector> &remote_iov_lists) = 0; }; #endif // __WORKER_H diff --git a/contrib/Dockerfile b/contrib/Dockerfile index 053eb87ad4..ba16ae0aaf 100644 --- a/contrib/Dockerfile +++ b/contrib/Dockerfile @@ -15,21 +15,29 @@ ARG BASE_IMAGE="nvcr.io/nvidia/cuda-dl-base" ARG BASE_IMAGE_TAG="25.03-cuda12.8-devel-ubuntu24.04" +ARG OS FROM ${BASE_IMAGE}:${BASE_IMAGE_TAG} +# Set default OS if not provided +ARG OS=${OS:-ubuntu24} ARG ARCH="x86_64" ARG DEFAULT_PYTHON_VERSION="3.12" -ARG UCX_REF="v1.19.x" +ARG UCX_REF="v1.19.0" ARG UCX_PREFIX="/usr" ARG UCX_PLUGIN_DIR="$UCX_PREFIX/lib/ucx" ARG NIXL_PREFIX="/usr/local/nixl" ARG NIXL_PLUGIN_DIR="$NIXL_PREFIX/lib/$ARCH-linux-gnu/plugins" +ARG NPROC +ARG WHL_DEFAULT_PYTHON_VERSIONS="3.12" +ARG EFA_INSTALLER_VERSION="latest" +ARG EFA_INSTALL_PATH="/opt/amazon/efa" RUN apt-get update -y && \ + apt-get install -y ubuntu-keyring && \ + apt-get update -y && \ DEBIAN_FRONTEND=noninteractive apt-get -y install \ ninja-build \ - pybind11-dev \ libclang-dev \ cmake \ libgflags-dev \ @@ -49,6 +57,10 @@ RUN apt-get update -y && \ flex \ libgtest-dev \ build-essential \ + python3.12-dev \ + clang \ + hwloc \ + libhwloc-dev \ libcurl4-openssl-dev libssl-dev uuid-dev zlib1g-dev # aws-sdk-cpp dependencies RUN DEBIAN_FRONTEND=noninteractive apt-get -y install \ @@ -56,16 +68,22 @@ RUN DEBIAN_FRONTEND=noninteractive apt-get -y install \ libnuma-dev librdmacm-dev ibverbs-providers WORKDIR /workspace -RUN git clone https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ +RUN git clone --depth 1 https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ cd etcd-cpp-apiv3 && \ sed -i '/^find_dependency(cpprestsdk)$/d' etcd-cpp-api-config.in.cmake && \ mkdir build && cd build && \ - cmake .. -DBUILD_ETCD_CORE_ONLY=ON -DCMAKE_BUILD_TYPE=Release && make -j$(nproc) && make install + cmake .. -DBUILD_ETCD_CORE_ONLY=ON -DCMAKE_BUILD_TYPE=Release && make -j${NPROC:-$(nproc)} && make install -RUN git clone --recurse-submodules https://github.com/aws/aws-sdk-cpp.git --branch 1.11.581 && \ +# Install EFA (Elastic Fabric Adapter) +RUN curl -fsSL "https://efa-installer.amazonaws.com/aws-efa-installer-${EFA_INSTALLER_VERSION}.tar.gz" | tar xz && \ + cd aws-efa-installer && \ + ./efa_installer.sh -y -g --skip-kmod --skip-limit-conf --no-verify && \ + ldconfig + +RUN git clone --recurse-submodules --depth 1 --shallow-submodules https://github.com/aws/aws-sdk-cpp.git --branch 1.11.581 && \ mkdir aws_sdk_build && cd aws_sdk_build && \ cmake ../aws-sdk-cpp/ -DCMAKE_BUILD_TYPE=Release -DBUILD_ONLY="s3" -DENABLE_TESTING=OFF -DCMAKE_INSTALL_PREFIX=/usr/local && \ - make -j$(nproc) && make install + make -j${NPROC:-$(nproc)} && make install COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/ @@ -85,29 +103,17 @@ RUN wget --tries=3 --waitretry=5 \ rm rustup-init* && \ chmod -R a+w $RUSTUP_HOME $CARGO_HOME -# Add Mellanox repository and install packages -RUN ARCH_SUFFIX=$(if [ "${ARCH}" = "aarch64" ]; then echo "arm64-sbsa"; else echo "${ARCH}"; fi) && \ - export PKG_CONFIG_PATH="/opt/mellanox/doca/lib/${ARCH_SUFFIX}-linux-gnu/pkgconfig:/opt/mellanox/dpdk/lib/${ARCH_SUFFIX}-linux-gnu/pkgconfig:$PKG_CONFIG_PATH" && \ - curl -fsSL https://linux.mellanox.com/public/repo/doca/3.0.0/ubuntu24.04/${ARCH_SUFFIX}/GPG-KEY-Mellanox.pub | \ - gpg --dearmor | tee /usr/share/keyrings/mellanox-archive-keyring.gpg && \ - echo "deb [signed-by=/usr/share/keyrings/mellanox-archive-keyring.gpg] https://linux.mellanox.com/public/repo/doca/3.0.0/ubuntu24.04/${ARCH_SUFFIX} ./" | \ - tee /etc/apt/sources.list.d/mellanox.list && \ - DEBIAN_FRONTEND=noninteractive apt update -y && \ - apt install -y --no-install-recommends \ - mlnx-dpdk mlnx-dpdk-dev \ - doca-sdk-common doca-sdk-dma doca-sdk-dpdk-bridge \ - doca-sdk-eth doca-sdk-flow doca-sdk-rdma doca-all \ - doca-sdk-gpunetio libdoca-sdk-gpunetio-dev - -WORKDIR /workspace/nixl -COPY . /workspace/nixl +# Add DOCA repository and install packages +RUN ARCH_SUFFIX=$(if [ "${ARCH}" = "aarch64" ]; then echo "arm64"; else echo "amd64"; fi) && \ + MELLANOX_OS=$(if [ "${OS}" = "ubuntu22" ]; then echo "ubuntu2204"; elif [ "${OS}" = "ubuntu24" ]; then echo "ubuntu2404"; else echo "${OS}"; fi) && \ + wget https://www.mellanox.com/downloads/DOCA/DOCA_v3.1.0/host/doca-host_3.1.0-091000-25.07-${MELLANOX_OS}_${ARCH_SUFFIX}.deb && \ + dpkg -i doca-host_3.1.0-091000-25.07-${MELLANOX_OS}_${ARCH_SUFFIX}.deb && \ + apt-get update -ENV LD_LIBRARY_PATH=/usr/local/lib:$LD_LIBRARY_PATH - -ENV VIRTUAL_ENV=/workspace/nixl/.venv -RUN rm -rf $VIRTUAL_ENV && uv venv $VIRTUAL_ENV --python $DEFAULT_PYTHON_VERSION && \ - # pybind11 pip install needed for ubuntu 22.04 - uv pip install --upgrade meson pybind11 patchelf +RUN if [ "$OS" = "ubuntu24" ]; then \ + apt-get install -y --no-install-recommends \ + doca-sdk-gpunetio libdoca-sdk-gpunetio-dev libdoca-sdk-verbs-dev; \ + fi RUN rm -rf /usr/lib/ucx RUN rm -rf /opt/hpcx/ucx @@ -130,13 +136,27 @@ RUN cd /usr/local/src && \ --with-gdrcopy=/usr/local \ --with-efa \ --enable-mt && \ - make -j && \ - make -j install-strip && \ + make -j${NPROC:-$(nproc)} && \ + make -j${NPROC:-$(nproc)} install-strip && \ ldconfig +WORKDIR /workspace/nixl +COPY . /workspace/nixl + +ENV LD_LIBRARY_PATH=/usr/local/lib:$EFA_INSTALL_PATH/lib:$LD_LIBRARY_PATH + +ENV VIRTUAL_ENV=/workspace/nixl/.venv +RUN rm -rf $VIRTUAL_ENV && uv venv $VIRTUAL_ENV --python $DEFAULT_PYTHON_VERSION && \ + # pybind11 pip install needed for ubuntu 22.04 + uv pip install --upgrade "meson>=0.64.0" pybind11 patchelf + +# Install pybind11 via apt +RUN apt-get update && apt-get install -y --no-install-recommends pybind11-dev + +ENV NIXL_PREFIX=$NIXL_PREFIX RUN rm -rf build && \ mkdir build && \ - uv run meson setup build/ --prefix=$NIXL_PREFIX && \ + uv run meson setup -Dlibfabric_path=$EFA_INSTALL_PATH build/ --prefix=$NIXL_PREFIX && \ cd build && \ ninja && \ ninja install @@ -147,6 +167,7 @@ RUN echo "$NIXL_PREFIX/lib/$ARCH-linux-gnu" > /etc/ld.so.conf.d/nixl.conf && \ RUN cd src/bindings/rust && cargo build --release --locked +# Build wheel using the build-wheel.sh script for better UCX plugin bundling and library management RUN ./contrib/build-wheel.sh \ --python-version $DEFAULT_PYTHON_VERSION \ --platform manylinux_2_39_$ARCH \ diff --git a/contrib/Dockerfile.manylinux b/contrib/Dockerfile.manylinux index bcb176a629..507f0f12e6 100644 --- a/contrib/Dockerfile.manylinux +++ b/contrib/Dockerfile.manylinux @@ -20,7 +20,7 @@ FROM ${BASE_IMAGE}:${BASE_IMAGE_TAG} ARG DEFAULT_PYTHON_VERSION="3.12" ARG ARCH="x86_64" -ARG UCX_REF="v1.19.x" +ARG UCX_REF="v1.19.0" RUN yum groupinstall -y 'Development Tools' && \ dnf install -y almalinux-release-synergy && \ @@ -97,12 +97,31 @@ RUN git clone --recurse-submodules -b v1.73.0 --depth 1 --shallow-submodules htt ENV LD_LIBRARY_PATH=/usr/local/lib:/usr/local/lib64:$LD_LIBRARY_PATH RUN cd /workspace && \ - git clone https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ + git clone --depth 1 https://github.com/etcd-cpp-apiv3/etcd-cpp-apiv3.git && \ cd etcd-cpp-apiv3 && \ sed -i '/^find_dependency(cpprestsdk)$/d' etcd-cpp-api-config.in.cmake && \ mkdir build && cd build && \ cmake .. -DBUILD_ETCD_CORE_ONLY=ON -DCMAKE_BUILD_TYPE=Release && make -j$(nproc) && make install +# The base image libcurl is linked against openssl 1.x, so we need to build from source +# in order to use openssl 3.x. This is needed to build aws-sdk-cpp. +RUN wget https://curl.se/download/curl-8.5.0.tar.gz && \ + tar xzf curl-8.5.0.tar.gz && cd curl-8.5.0 && \ + ./configure --prefix=/usr/local --with-ssl=/usr/local/openssl3 --enable-shared && \ + make -j$(nproc) && make install + +RUN git clone --recurse-submodules --depth 1 --shallow-submodules https://github.com/aws/aws-sdk-cpp.git --branch 1.11.581 +RUN mkdir aws_sdk_build && cd aws_sdk_build && \ + export LDFLAGS="-L/usr/local/openssl3/lib64 -L/usr/local/openssl3/lib" && \ + export CFLAGS="-I/usr/local/openssl3/include" && \ + export CXXFLAGS="-I/usr/local/openssl3/include" && \ + cmake ../aws-sdk-cpp/ -DCMAKE_BUILD_TYPE=Release -DBUILD_ONLY="s3" -DENABLE_TESTING=OFF -DCMAKE_INSTALL_PREFIX=/usr/local \ + -DCMAKE_PREFIX_PATH="/usr/local/openssl3;/usr/local" \ + -DCURL_LIBRARY=/usr/local/lib/libcurl.so \ + -DCURL_INCLUDE_DIR=/usr/local/include \ + -DOPENSSL_USE_STATIC_LIBS=OFF && \ + make -j${NPROC:-$(nproc)} && make install + COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/ ENV RUSTUP_HOME=/usr/local/rustup \ @@ -123,6 +142,17 @@ RUN wget --tries=3 --waitretry=5 "https://static.rust-lang.org/rustup/archive/1. rm rustup-init && \ chmod -R a+w $RUSTUP_HOME $CARGO_HOME +RUN wget https://www.mellanox.com/downloads/DOCA/DOCA_v3.1.0/host/doca-host-3.1.0-091000_25.07_rhel89.${ARCH}.rpm && \ + rpm -i doca-host-3.1.0-091000_25.07_rhel89.${ARCH}.rpm && \ + dnf install -y libnl3-devel && \ + cd /usr/share/doca-host-3.1.0/repo/Packages/ && \ + rpm -ivh --nodeps doca-sdk-common-*rpm && \ + rpm -ivh --nodeps doca-sdk-rdma-*rpm && \ + rpm -ivh --nodeps doca-sdk-verbs-*rpm && \ + rpm -ivh --nodeps doca-sdk-gpunetio-*rpm && \ + # Check that gpunetio development package is installed correctly + pkg-config --cflags --libs doca-gpunetio + ENV LD_LIBRARY_PATH=/usr/local/lib:$LD_LIBRARY_PATH ENV CUDA_PATH=/usr/local/cuda @@ -136,12 +166,16 @@ RUN rm -rf /usr/lib/ucx RUN rm -rf /opt/hpcx/ucx RUN cd /workspace && \ - git clone https://github.com/NVIDIA/gdrcopy.git && \ - cd gdrcopy/packages && \ + git clone --depth 1 https://github.com/NVIDIA/gdrcopy.git && \ + cd gdrcopy && \ + git fetch --tags --depth=1 && \ + latest_tag=$(git describe --tags "$(git rev-list --tags --max-count=1)") && \ + git checkout "$latest_tag" && \ + cd packages && \ CUDA=/usr/local/cuda ./build-rpm-packages.sh && \ - rpm -Uvh gdrcopy-kmod-2.5-1dkms.el8.noarch.rpm && \ - rpm -Uvh gdrcopy-2.5-1.el8.$ARCH.rpm && \ - rpm -Uvh gdrcopy-devel-2.5-1.el8.noarch.rpm + rpm -Uvh gdrcopy-kmod-*.el8.noarch.rpm && \ + rpm -Uvh gdrcopy-*.el8.$ARCH.rpm && \ + rpm -Uvh gdrcopy-devel-*.el8.noarch.rpm RUN cd /usr/local/src && \ git clone https://github.com/openucx/ucx.git && \ diff --git a/contrib/aws-efa/aws_test.sh b/contrib/aws-efa/aws_test.sh index e0ab70b975..d91c404211 100755 --- a/contrib/aws-efa/aws_test.sh +++ b/contrib/aws-efa/aws_test.sh @@ -31,6 +31,7 @@ usage() { echo "" echo "Optional environment variables:" echo " CONTAINER_IMAGE - Container image to use (default: nvcr.io/nvidia/pytorch:25.02-py3)" + echo " TEST_TIMEOUT - Timeout for test execution in minutes" exit 1 } @@ -64,6 +65,12 @@ setup_cmd="set -x && \ cd nixl && \ ${GIT_CHECKOUT_CMD}" build_cmd=".gitlab/build.sh \${NIXL_INSTALL_DIR} \${UCX_INSTALL_DIR}" + +# Add timeout only if TEST_TIMEOUT is set (expects minutes) +if [ -n "$TEST_TIMEOUT" ]; then + test_cmd="timeout ${TEST_TIMEOUT}m ${test_cmd}" +fi + export AWS_CMD="${setup_cmd} && ${build_cmd} && ${test_cmd}" # Generate AWS job properties json from template diff --git a/contrib/aws-efa/aws_vars.template b/contrib/aws-efa/aws_vars.template index 45963666cc..7a2da6511b 100644 --- a/contrib/aws-efa/aws_vars.template +++ b/contrib/aws-efa/aws_vars.template @@ -9,6 +9,10 @@ ], "image": "${CONTAINER_IMAGE}", "env": [ + { + "name": "UCX_VERSION", + "value": "${UCX_VERSION}" + }, { "name": "UCX_INSTALL_DIR", "value": "/usr/local" diff --git a/contrib/build-container.sh b/contrib/build-container.sh index fbc94846e2..cc647b5fcc 100755 --- a/contrib/build-container.sh +++ b/contrib/build-container.sh @@ -35,8 +35,9 @@ ARCH=$(uname -m) WHL_BASE=manylinux_2_39 WHL_PLATFORM=${WHL_BASE}_${ARCH} WHL_PYTHON_VERSIONS="3.12" -UCX_REF=v1.19.x +UCX_REF=${UCX_REF:-v1.19.0} OS="ubuntu24" +NPROC=${NPROC:-$(nproc)} get_options() { while :; do @@ -193,6 +194,8 @@ BUILD_ARGS+=" --build-arg WHL_PYTHON_VERSIONS=$WHL_PYTHON_VERSIONS" BUILD_ARGS+=" --build-arg WHL_PLATFORM=$WHL_PLATFORM" BUILD_ARGS+=" --build-arg ARCH=$ARCH" BUILD_ARGS+=" --build-arg UCX_REF=$UCX_REF" +BUILD_ARGS+=" --build-arg NPROC=$NPROC" +BUILD_ARGS+=" --build-arg OS=$OS" show_build_options diff --git a/contrib/build-wheel.sh b/contrib/build-wheel.sh index 78f876a1ed..fa313d49a1 100755 --- a/contrib/build-wheel.sh +++ b/contrib/build-wheel.sh @@ -86,7 +86,7 @@ uv build --wheel --out-dir $TMP_DIR --python $PYTHON_VERSION # Bundle libraries uv pip install auditwheel patchelf -uv run auditwheel repair --exclude libcuda.so.1 --exclude 'libssl*' --exclude 'libcrypto*' $TMP_DIR/nixl-*.whl --plat $WHL_PLATFORM --wheel-dir $OUTPUT_DIR +uv run auditwheel repair --exclude 'libcuda*' --exclude 'libcufile*' --exclude 'libssl*' --exclude 'libcrypto*' $TMP_DIR/nixl-*.whl --plat $WHL_PLATFORM --wheel-dir $OUTPUT_DIR uv run ./contrib/wheel_add_ucx_plugins.py --ucx-plugins-dir $UCX_PLUGINS_DIR --nixl-plugins-dir $NIXL_PLUGINS_DIR $OUTPUT_DIR/*.whl diff --git a/contrib/wheel_add_ucx_plugins.py b/contrib/wheel_add_ucx_plugins.py index 124ff81ce3..d9bd8d0698 100755 --- a/contrib/wheel_add_ucx_plugins.py +++ b/contrib/wheel_add_ucx_plugins.py @@ -19,11 +19,14 @@ import base64 import csv import hashlib +import logging import os import shutil import tempfile import zipfile +logger = logging.getLogger(__name__) + def extract_wheel(wheel_path): """ @@ -32,7 +35,7 @@ def extract_wheel(wheel_path): Path to the temporary directory. The caller is responsible for cleaning up the directory. """ temp_dir = tempfile.mkdtemp() - print(f"Extracting wheel {wheel_path} to {temp_dir}") + logger.info("Extracting wheel %s to %s", wheel_path, temp_dir) with zipfile.ZipFile(wheel_path, "r") as zip_ref: zip_ref.extractall(temp_dir) return temp_dir @@ -82,7 +85,7 @@ def create_wheel(wheel_path, temp_dir): """ Create a wheel from a temporary directory. """ - print(f"Creating wheel {wheel_path} from {temp_dir}") + logger.info("Creating wheel %s from %s", wheel_path, temp_dir) update_wheel_record_file(temp_dir) with zipfile.ZipFile( wheel_path, "w", compression=zipfile.ZIP_DEFLATED, compresslevel=9 @@ -190,7 +193,7 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): if not os.path.exists(pkg_libs_dir): raise FileNotFoundError(f"nixl.libs directory not found in wheel: {wheel_path}") - print("Listing existing libs:") + logger.debug("Listing existing libs:") name_map = get_repaired_lib_name_map(pkg_libs_dir) # Ensure that all of them in name_map have RPATH set to $ORIGIN @@ -203,20 +206,20 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): rpath = "$ORIGIN" else: rpath = "$ORIGIN:" + rpath - print(f"Setting rpath for {fpath} to {rpath}") + logger.debug("Setting rpath for %s to %s", fpath, rpath) ret = os.system(f"patchelf --set-rpath '{rpath}' {fpath}") if ret != 0: raise RuntimeError(f"Failed to set rpath for {fpath}") pkg_plugins_dir = os.path.join(pkg_libs_dir, install_dirname) - print(f"Copying plugins from {sys_plugins_dir} to {pkg_plugins_dir}") + logger.debug("Copying plugins from %s to %s", sys_plugins_dir, pkg_plugins_dir) copied_files = copytree(sys_plugins_dir, pkg_plugins_dir) if not copied_files: raise RuntimeError(f"No plugins found in {sys_plugins_dir}") # Patch all libs to load plugin deps from the wheel for fname in copied_files: - print(f"Patching {fname}") + logger.debug("Patching %s", fname) fpath = os.path.join(pkg_plugins_dir, fname) if os.path.isfile(fpath) and ".so" in fname: rpath = os.popen(f"patchelf --print-rpath {fpath}").read().strip() @@ -224,7 +227,7 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): rpath = "$ORIGIN/..:$ORIGIN" else: rpath = "$ORIGIN/..:$ORIGIN:" + rpath - print(f"Setting rpath for {fpath} to {rpath}") + logger.debug("Setting rpath for %s to %s", fpath, rpath) ret = os.system(f"patchelf --set-rpath '{rpath}' {fpath}") if ret != 0: raise RuntimeError(f"Failed to set rpath for {fpath}") @@ -234,7 +237,9 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): base_name = libname.split(".")[0] if base_name in name_map: packaged_name = name_map[base_name] - print(f"Replacing {libname} with {packaged_name} in {fpath}") + logger.debug( + "Replacing %s with %s in %s", libname, packaged_name, fpath + ) ret = os.system( f"patchelf --replace-needed {libname} {packaged_name} {fpath}" ) @@ -243,7 +248,7 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): f"Failed to replace {libname} with {packaged_name} in {fpath}" ) # Check that there is no breakage introduced in the patched lib - print(f"Checking that {fpath} loads") + logger.debug("Checking that %s loads", fpath) original_deps = get_lib_deps(os.path.join(sys_plugins_dir, fname)) for libname, libpath in get_lib_deps(fpath).items(): if libpath is None: @@ -255,7 +260,7 @@ def add_plugins(wheel_path, sys_plugins_dir, install_dirname): create_wheel(wheel_path, temp_dir) shutil.rmtree(temp_dir) - print(f"Added plugins to wheel: {wheel_path}") + logger.info("Added plugins to wheel: %s", wheel_path) def main(): diff --git a/docs/BackendGuide.md b/docs/BackendGuide.md index b14f3d61e8..b68b286a85 100644 --- a/docs/BackendGuide.md +++ b/docs/BackendGuide.md @@ -51,12 +51,10 @@ The key/value parameters are a map of strings to byte arrays that are passed fro * supportsLocal(): Indicates if the backend supports transfers within a node * supportsRemote(): Indicates if the backend supports transfers across nodes * supportsNotif(): Indicates if the backend supports notifications -* supportsProgressThread(): Indicates if the backend supports progress() method. That method should call the underlying procedure of progressing transfers for this backend. * getSupportedMems(): Indicates memory types supported by the backend -Based on the first 4 methods (supports*), the required methods to be implemented change. For instance, UCX backend implements all as it supports all scenarios, while GDS backend only has supportsLocal, detailed more in Example implementations. Note that a network backend should have supportsRemote and supportsNotif to be set to true, and preferably supportsLocal also to true, so another backend doesn’t need to be involved for local transfers. For a storage backend, it should have supportsLocal and supportsNotif is optional. supportsProgressThread is optional for both cases. Additionally, a backend that supportsRemote must also support supportNotifs. +Based on the first 3 methods (supports*), the required methods to be implemented change. For instance, UCX backend implements all as it supports all scenarios, while GDS backend only has supportsLocal, detailed more in Example implementations. Note that a network backend must have supportsRemote and supportsNotif to be set to true, and preferably supportsLocal also to true, so another backend doesn't need to be involved for local transfers. For a storage backend, it should have supportsLocal and supportsNotif is optional. -Note that supportProgressThread is an indicator whether a backend has implemented the progress() method, but does not imply how the progress thread is implemented. During creation of a backend, the provided init params indicate how the progress thread is intended to be used. For instance, if the enablement of progress thread is set to false, while a backend cannot work without a separate progress thread, the backend creation would fail. This flag is useful for the NIXL agent if we want to provide some agent level guarantees, such as minimum time between calls to progress for backends, or if a central progress method is implemented (for future proofing, not currently implemented). ### Connection Management: * connect(): Initiates connection to a remote agent. @@ -106,12 +104,6 @@ Finally, note that a call to releaseXferReq should not block and be asynchronous Note that getNotif does not know which agent it should look for to receive the notification. So there should be a method to extract the agent name from the notification received, corresponding to a transfer. genNotif generates a notification which is not bound to any transfers, and does not provide any ordering guarantees. If a backend does not set supportsNotifications, these two methods are not needed. -### Progress Thread: - -* progress(): Makes progress on transfers and notifications. - -If a backend requires a progress call, such as UCX, to proceed with the transfers, for both check of transfer status or received notification, they can implement a progress thread, and a frequency of waking up that thread will be passed during backend creation. In addition, each time a user calls to check a transfer status, or check received notifications, this method is called, enabling progress if a progress thread is not implemented. - ## Descriptor List Abstraction A key underlying abstraction for NIXL library is a descriptor list, that is made of a memory space (host/GPU/block/File/Obj-Store) and a list of descriptors. There are 2 types of descriptors used for the SB API. @@ -142,9 +134,9 @@ The plugin manager maintains API versioning of these above APIs. This can allow ## Comparing two plugins as an example -NIXL UCX plugin provides networking across different nodes, while GDS plugin provides storage access. Moreover, UCX plugin sets all of the “supports” flags, while GDS only has the supportsLocal flag set. The reason being UCX requires a progress thread and provides notifications, and can do transfers within an Agent, for instance from GPU to CPU, and across Agents. Therefore, it should implement all of the methods mentioned previously. +NIXL UCX plugin provides networking across different nodes, while GDS plugin provides storage access. UCX plugin sets all of the “supports” flags, while GDS only has the supportsLocal flag set. The reason being UCX is a network plugin that should support inter-agent communication and notifications, and it also supports intra-agent transfers, for instance from GPU to CPU. -However, for NIXL storage backends, there is no need to run a NIXL agent on a remote storage node. Instead, a distributed storage client on the local agent talks to the remote distributed storage, and therefore from NIXL agent point of view for all storage, whether local or remote, it has to talk to this local storage client. In other words, all the transfers are loopback to the agent itself. For the current use case, there is no need for notifications within the same agent, or a progress thread either. +However, for NIXL storage backends, there is no need to run a NIXL agent on a remote storage node. Instead, a distributed storage client on the local agent talks to the remote distributed storage, and therefore from NIXL agent point of view for all storage, whether local or remote, it has to talk to this local storage client. In other words, all the transfers are loopback to the agent itself. For the current use case, there is no need for notifications within the same agent. Moreover, the GDS plugin does not require a local connection to itself, so it returns SUCCESS for connect and disconnect, and for loadLocal simply returns back the input pointer as its output. The only 6 remaining methods that it has to implement are: @@ -213,7 +205,7 @@ Note that inside a transfer, a backend might provide methods for network resilie ### Get transfer status: -The agent will call the backend specific transfer handle that is stored within the agent transfer handle, and check the status of the transfer. This is achieved through a call to **checkXfer** in the SB API. Internal to the backend, they can call the **progress** method in SB API, if that’s necessary to get the latest status of the transfers. If the agent is run in progress thread mode, the agent will call that periodically, and therefore reduce the load on this internal call. +The agent will call the backend specific transfer handle that is stored within the agent transfer handle, and check the status of the transfer. This is achieved through a call to **checkXfer** in the SB API. Internal to the backend, they can call their internal progress method, if that’s necessary to get the latest status of the transfers. ### Invalidate transfer request: @@ -221,7 +213,7 @@ The agent will call the **releaseReqH** from the SB API on the backend specific ### Get notifications: -The agent will iterate over all the backends that support notification, and call their **getNotifs** from the SB API, which will return a list of notifications received from each remote node between the previous call to this method and this time. Then the agent will merge the results from all such backends, and append them to the map that the user has provided. Similar to get transfer status, Internal to the backend, they can call the **progress** method in SB API, if that’s necessary to get the latest notifications received from the transfers initiated by the other agents towards them. If the agent is run in progress thread mode, the agent will call that periodically, and therefore reduce the load on this internal call. +The agent will iterate over all the backends that support notification, and call their **getNotifs** from the SB API, which will return a list of notifications received from each remote node between the previous call to this method and this time. Then the agent will merge the results from all such backends, and append them to the map that the user has provided. Similar to get transfer status, Internal to the backend, they can call their internal progress method, if that’s necessary to get the latest notifications received from the transfers initiated by the other agents towards them. ### Generate notification: @@ -229,4 +221,4 @@ If a backend is provided by the user, the agent will call **genNotif** from the ### Destructor: -When an agent is getting destroyed at the end of the application, it will deregister all the remaining memories that were not deregistered by the application (bad practice, but agent takes care of it). Then for each of the backends it will call their **destructor** from the SB API, and finally do the rest of internal clean up. \ No newline at end of file +When an agent is getting destroyed at the end of the application, it will deregister all the remaining memories that were not deregistered by the application (bad practice, but agent takes care of it). Then for each of the backends it will call their **destructor** from the SB API, and finally do the rest of internal clean up. diff --git a/docs/figures/nixl_high_level.png b/docs/figures/nixl_high_level.png index 6b91048155..7ab708893a 100644 Binary files a/docs/figures/nixl_high_level.png and b/docs/figures/nixl_high_level.png differ diff --git a/docs/figures/nixl_sb_api.png b/docs/figures/nixl_sb_api.png index 63735dd328..503c28b177 100644 Binary files a/docs/figures/nixl_sb_api.png and b/docs/figures/nixl_sb_api.png differ diff --git a/docs/telemetry.md b/docs/telemetry.md new file mode 100644 index 0000000000..23f04e630c --- /dev/null +++ b/docs/telemetry.md @@ -0,0 +1,100 @@ +# NIXL Telemetry System + +## Overview + +The NIXL telemetry system provides real-time monitoring and performance tracking capabilities for NIXL applications. It collects various metrics and events during runtime and stores them in shared memory buffers that can be read by telemetry reader applications. + +## Architecture + +### Telemetry Components + +1. **Telemetry Collection**: Built into the NIXL core library, collects events and metrics +2. **Shared Memory Buffer**: Cyclic buffer implementation for efficient event storage +3. **Telemetry Readers**: C++ and Python applications to read and display telemetry data + +### Event Structure + +Each telemetry event contains: +- **Timestamp**: Microsecond precision timestamp +- **Category**: Event category for filtering and aggregation +- **Event Name**: Descriptive name/identifier for the event +- **Value**: Numeric value associated with the event + +### Event Categories + +The telemetry system supports the following event categories: + +| Category | Description | Example Events | +|----------|-------------|----------------| +| `NIXL_TELEMETRY_MEMORY` | Memory operations | Memory registration, deregistration, allocation | +| `NIXL_TELEMETRY_TRANSFER` | Data transfer operations | Bytes transmitted/received, request counts | +| `NIXL_TELEMETRY_CONNECTION` | Connection management | Connect, disconnect events | +| `NIXL_TELEMETRY_BACKEND` | Backend-specific operations | Backend initialization, configuration | +| `NIXL_TELEMETRY_ERROR` | Error events | Error counts by type | +| `NIXL_TELEMETRY_PERFORMANCE` | Performance metrics | Transaction times, latency measurements | +| `NIXL_TELEMETRY_SYSTEM` | System-level events | Process start/stop, resource usage | +| `NIXL_TELEMETRY_CUSTOM` | Custom/user-defined events | Application-specific metrics | + +## Enabling Telemetry + +### Runtime Configuration + +Telemetry is controlled by environment variables: + +| Variable | Description | Default | +|----------|-------------|---------| +| `NIXL_TELEMETRY_ENABLE` | Enable telemetry collection | Disabled | +| `NIXL_TELEMETRY_DIR` | Directory for telemetry files | - | +| `NIXL_TELEMETRY_BUFFER_SIZE` | Number of events in buffer | `4096` | +| `NIXL_TELEMETRY_RUN_INTERVAL` | Flush interval (ms) | `100` | + +- NIXL_TELEMETRY_ENABLE can be set to y/yes/on/1 to be enabled, and n/no/off/0 (or not set) to be disabled, +- If NIXL_TELEMETRY_ENABLE is set to enabled but NIXL_TELEMETRY_DIR is not set, no telemetry file is generated and NIXL_TELEMETRY_RUN_INTERVAL is not used. + +## Telemetry File Format + +Telemetry data is stored in shared memory files with the agent name passed when creating the agent. + +## Using Telemetry Readers + +### C++ Telemetry Reader + +The C++ telemetry reader (`telemetry_reader.cpp`) provides a robust way to read and display telemetry events. + +#### Running the C++ Reader + +```bash +# Read from a specific telemetry file +./builddir/examples/cpp/telemetry_reader /tmp/agent_name +``` + +### Python Telemetry Reader + +The Python telemetry reader (`telemetry_reader.py`) provides similar functionality with additional features. + +#### Running the Python Reader + +```bash +# Read from a specific telemetry file +python3 examples/python/telemetry_reader.py --telemetry_path /tmp/agent_name +``` + +## Example Output + +Both readers produce similar formatted output: + +``` +=== NIXL Telemetry Event === +Timestamp: 2025-01-15 14:30:25.123456 +Category: TRANSFER +Event: agent_tx_bytes +Value: 1048576 +=========================== + +=== NIXL Telemetry Event === +Timestamp: 2025-01-15 14:30:25.124567 +Category: MEMORY +Event: agent_memory_registered +Value: 4096 +=========================== +``` diff --git a/examples/cpp/meson.build b/examples/cpp/meson.build index 7655d6acf2..e0bc94b701 100644 --- a/examples/cpp/meson.build +++ b/examples/cpp/meson.build @@ -28,3 +28,9 @@ if etcd_dep.found() link_with: [serdes_lib], install: true) endif + +telemetry_reader = executable('telemetry_reader', + 'telemetry_reader.cpp', + dependencies: [nixl_dep, nixl_common_deps], + include_directories: [nixl_inc_dirs, utils_inc_dirs], + install: true) diff --git a/examples/cpp/telemetry_reader.cpp b/examples/cpp/telemetry_reader.cpp new file mode 100644 index 0000000000..34e5c0e9d4 --- /dev/null +++ b/examples/cpp/telemetry_reader.cpp @@ -0,0 +1,147 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + + +namespace fs = std::filesystem; + +#include "common/cyclic_buffer.h" +#include "telemetry_event.h" + +volatile sig_atomic_t g_running = true; + +// Signal handler for Ctrl+C +void +signal_handler(int signal) { + if (signal == SIGINT) { + g_running = false; + } +} + +std::string +format_timestamp(uint64_t timestamp_us) { + auto time_point = + std::chrono::system_clock::time_point(std::chrono::microseconds(timestamp_us)); + auto time_t = std::chrono::system_clock::to_time_t(time_point); + + std::stringstream ss; + ss << std::put_time(std::localtime(&time_t), "%Y-%m-%d %H:%M:%S"); + + auto microseconds = timestamp_us % 1000000; + ss << "." << std::setfill('0') << std::setw(6) << microseconds; + + return ss.str(); +} + +std::string +format_bytes(uint64_t bytes) { + const char *units[] = {"B", "KB", "MB", "GB", "TB"}; + int unit_index = 0; + double value = static_cast(bytes); + + while (value >= 1024.0 && unit_index < 4) { + value /= 1024.0; + unit_index++; + } + + std::stringstream ss; + ss << std::fixed << std::setprecision(2) << value << " " << units[unit_index]; + return ss.str(); +} + +void +print_telemetry_event(const nixlTelemetryEvent &event) { + // Can be extended to more general ostream if needed + // friend std::ostream &operator<<(std::ostream &os, const nixlTelemetryEvent &event) + std::cout << "\n=== NIXL Telemetry Event ===" << std::endl; + std::cout << "Timestamp: " << format_timestamp(event.timestampUs_) << std::endl; + std::cout << "Category: " << nixlEnumStrings::telemetryCategoryStr(event.category_) + << std::endl; + std::cout << "Event name: " << event.eventName_ << std::endl; + std::cout << "Value: " << event.value_ << std::endl; + + std::cout << "===========================" << std::endl; +} + +void +usage() { + std::cout << "Usage: telemetry_reader " << std::endl; + std::cout << "Options:" << std::endl; + std::cout << " Path to the telemetry file" << std::endl; + exit(0); +} + +int +main(int argc, char *argv[]) { + if (argc < 2 || argv[1] == std::string("-h") || argv[1] == std::string("--help")) { + usage(); + } + + std::cout << "Telemetry path: " << argv[1] << std::endl; + auto telemetry_path = argv[1]; + + if (!fs::exists(telemetry_path)) { + std::cerr << "Telemetry file " << telemetry_path << " does not exist" << std::endl; + return 1; + } + + signal(SIGINT, signal_handler); + + try { + std::cout << "Opening telemetry buffer: " << telemetry_path << std::endl; + std::cout << "Press Ctrl+C to stop reading telemetry..." << std::endl; + + sharedRingBuffer buffer(telemetry_path, false, TELEMETRY_VERSION); + + std::cout << "Successfully opened telemetry buffer (version: " << buffer.version() << ")" + << std::endl; + std::cout << "Buffer capacity: " << buffer.capacity() << " events" << std::endl; + + nixlTelemetryEvent event; + uint64_t event_count = 0; + + while (g_running) { + if (buffer.pop(event)) { + event_count++; + print_telemetry_event(event); + } else { + std::this_thread::sleep_for(std::chrono::milliseconds(100)); + } + } + + std::cout << "\nTotal events read: " << event_count << std::endl; + std::cout << "Final buffer size: " << buffer.size() << " events" << std::endl; + } + catch (const std::exception &e) { + std::cerr << "Error: " << e.what() << std::endl; + return 1; + } + + return 0; +} diff --git a/examples/python/blocking_send_recv_example.py b/examples/python/blocking_send_recv_example.py index 878d145ed9..da4064c535 100755 --- a/examples/python/blocking_send_recv_example.py +++ b/examples/python/blocking_send_recv_example.py @@ -20,6 +20,9 @@ import torch from nixl._api import nixl_agent, nixl_agent_config +from nixl.logging import get_logger + +logger = get_logger(__name__) def parse_args(): @@ -58,11 +61,11 @@ def parse_args(): else: tensors = [torch.zeros(10, dtype=torch.float32) for _ in range(2)] - print(f"{args.mode} Tensors: {tensors}") + logger.info("Running test with %s tensors in mode %s", tensors, args.mode) reg_descs = agent.register_memory(tensors) if not reg_descs: # Same as reg_descs if successful - print("Memory registration failed.") + logger.error("Memory registration failed.") exit() # Target code @@ -78,7 +81,7 @@ def parse_args(): agent.send_notif("initiator", target_desc_str) - print("Waiting for transfer") + logger.info("Waiting for transfer") # Waiting for transfer # For now the notification is just UUID, could be any python bytes. @@ -88,7 +91,7 @@ def parse_args(): continue # Initiator code else: - print("Initiator sending to " + args.ip) + logger.info("Initiator sending to %s", args.ip) agent.fetch_remote_metadata("target", args.ip, args.port) agent.send_local_metadata(args.ip, args.port) @@ -105,24 +108,24 @@ def parse_args(): while not ready: ready = agent.check_remote_metadata("target") - print("Ready for transfer") + logger.info("Ready for transfer") xfer_handle = agent.initialize_xfer( "READ", initiator_descs, target_descs, "target", "UUID" ) if not xfer_handle: - print("Creating transfer failed.") + logger.error("Creating transfer failed.") exit() state = agent.transfer(xfer_handle) if state == "ERR": - print("Posting transfer failed.") + logger.error("Posting transfer failed.") exit() while True: state = agent.check_xfer_state(xfer_handle) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": break @@ -130,9 +133,9 @@ def parse_args(): # Verify data after read for i, tensor in enumerate(tensors): if not torch.allclose(tensor, torch.ones(10)): - print(f"Data verification failed for tensor {i}.") + logger.error("Data verification failed for tensor %d.", i) exit() - print(f"{args.mode} Data verification passed - {tensors}") + logger.info("%s Data verification passed", args.mode) if args.mode != "target": agent.remove_remote_agent("target") @@ -141,4 +144,4 @@ def parse_args(): agent.deregister_memory(reg_descs) - print("Test Complete.") + logger.info("Test Complete.") diff --git a/examples/python/nixl_api_example.py b/examples/python/nixl_api_example.py index 316e84378d..13333389d4 100755 --- a/examples/python/nixl_api_example.py +++ b/examples/python/nixl_api_example.py @@ -22,13 +22,17 @@ import nixl._utils as nixl_utils from nixl._api import nixl_agent, nixl_agent_config +from nixl.logging import get_logger + +# Configure logging +logger = get_logger(__name__) + if __name__ == "__main__": buf_size = 256 # Allocate memory and register with NIXL - print("Using NIXL Plugins from:") - print(os.environ["NIXL_PLUGIN_DIR"]) + logger.info("Using NIXL Plugins from:\n%s", os.environ["NIXL_PLUGIN_DIR"]) # Example using nixl_agent_config agent_config = nixl_agent_config(backends=["UCX"]) @@ -37,14 +41,17 @@ plugin_list = nixl_agent1.get_plugin_list() assert "UCX" in plugin_list - print("Plugin parameters") - print(nixl_agent1.get_plugin_mem_types("UCX")) - print(nixl_agent1.get_plugin_params("UCX")) + logger.info( + "Plugin parameters:\n%s\n%s", + nixl_agent1.get_plugin_mem_types("UCX"), + nixl_agent1.get_plugin_params("UCX"), + ) - print("\nLoaded backend parameters") - print(nixl_agent1.get_backend_mem_types("UCX")) - print(nixl_agent1.get_backend_params("UCX")) - print() + logger.info( + "Backend parameters:\n%s\n%s", + nixl_agent1.get_backend_mem_types("UCX"), + nixl_agent1.get_backend_params("UCX"), + ) addr1 = nixl_utils.malloc_passthru(buf_size * 2) addr2 = addr1 + buf_size @@ -52,21 +59,19 @@ agent1_addrs = [(addr1, buf_size, 0), (addr2, buf_size, 0)] agent1_strings = [(addr1, buf_size, 0, "a"), (addr2, buf_size, 0, "b")] - agent1_reg_descs = nixl_agent1.get_reg_descs(agent1_strings, "DRAM", is_sorted=True) - agent1_xfer_descs = nixl_agent1.get_xfer_descs(agent1_addrs, "DRAM", is_sorted=True) + agent1_reg_descs = nixl_agent1.get_reg_descs(agent1_strings, "DRAM") + agent1_xfer_descs = nixl_agent1.get_xfer_descs(agent1_addrs, "DRAM") # Prefer numpy arrays for performance agent1_addrs_np = np.array(agent1_addrs) - agent1_xfer_descs_np = nixl_agent1.get_xfer_descs( - agent1_addrs_np, "DRAM", is_sorted=True - ) - agent1_reg_descs_np = nixl_agent1.get_reg_descs( - agent1_addrs_np, "DRAM", is_sorted=True - ) + agent1_xfer_descs_np = nixl_agent1.get_xfer_descs(agent1_addrs_np, "DRAM") + agent1_reg_descs_np = nixl_agent1.get_reg_descs(agent1_addrs_np, "DRAM") assert agent1_xfer_descs == agent1_xfer_descs_np assert agent1_reg_descs == agent1_reg_descs_np - print(agent1_reg_descs, agent1_reg_descs_np) + logger.debug( + "Registration descriptors: %s %s", agent1_reg_descs, agent1_reg_descs_np + ) # Just for tensor test tensors = [torch.zeros(10, dtype=torch.float32) for _ in range(2)] @@ -83,16 +88,16 @@ agent2_addrs = [(addr3, buf_size, 0), (addr4, buf_size, 0)] agent2_strings = [(addr3, buf_size, 0, "a"), (addr4, buf_size, 0, "b")] - agent2_reg_descs = nixl_agent2.get_reg_descs(agent2_strings, "DRAM", is_sorted=True) - agent2_xfer_descs = nixl_agent2.get_xfer_descs(agent2_addrs, "DRAM", is_sorted=True) + agent2_reg_descs = nixl_agent2.get_reg_descs(agent2_strings, "DRAM") + agent2_xfer_descs = nixl_agent2.get_xfer_descs(agent2_addrs, "DRAM") - agent2_descs = nixl_agent2.register_memory(agent2_reg_descs, is_sorted=True) + agent2_descs = nixl_agent2.register_memory(agent2_reg_descs) assert agent2_descs is not None # Exchange metadata meta = nixl_agent1.get_agent_metadata() remote_name = nixl_agent2.add_remote_agent(meta) - print("Loaded name from metadata:", remote_name, flush=True) + logger.info("Loaded name from metadata: %s", remote_name) serdes = nixl_agent1.get_serialized_descs(agent1_reg_descs) src_descs_recvd = nixl_agent2.deserialize_descs(serdes) @@ -103,7 +108,7 @@ "READ", agent2_xfer_descs, agent1_xfer_descs, remote_name, b"UUID1" ) if not xfer_handle_1: - print("Creating transfer failed.") + logger.error("Creating transfer failed.") exit() # test multiple postings @@ -118,20 +123,20 @@ if not init_done: state = nixl_agent2.check_xfer_state(xfer_handle_1) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": init_done = True - print("Initiator done") + logger.info("Initiator done") if not target_done: if nixl_agent1.check_remote_xfer_done("initiator", b"UUID1"): target_done = True - print("Target done") + logger.info("Target done") # prep transfer mode local_prep_handle = nixl_agent2.prep_xfer_dlist( - "NIXL_INIT_AGENT", [(addr3, buf_size, 0), (addr4, buf_size, 0)], "DRAM", True + "NIXL_INIT_AGENT", [(addr3, buf_size, 0), (addr4, buf_size, 0)], "DRAM" ) remote_prep_handle = nixl_agent2.prep_xfer_dlist( remote_name, agent1_xfer_descs, "DRAM" @@ -145,30 +150,29 @@ test_notif = str.encode("DESCS: ") + serdes nixl_agent2.send_notif(remote_name, test_notif) - print("sent notif ") - print(test_notif) + logger.info("sent notif: \n%s", test_notif) notif_recv = False while not notif_recv: notif_map = nixl_agent1.get_new_notifs() if "initiator" in notif_map: - print("received message from initiator") + logger.info("received message from initiator") for msg in notif_map["initiator"]: if msg == test_notif: notif_recv = True - print("notif test complete, doing transfer 2\n") + logger.info("notif test complete, doing transfer 2") xfer_handle_2 = nixl_agent2.make_prepped_xfer( "WRITE", local_prep_handle, [0, 1], remote_prep_handle, [1, 0], b"UUID2" ) if not local_prep_handle or not remote_prep_handle: - print("Preparing transfer side handles failed.") + logger.error("Preparing transfer side handles failed.") exit() if not xfer_handle_2: - print("Make prepped transfer failed.") + logger.error("Make prepped transfer failed.") exit() state = nixl_agent2.transfer(xfer_handle_2) @@ -177,22 +181,22 @@ target_done = False init_done = False - print("Transfer 2 started") + logger.info("Transfer 2 started") while (not init_done) or (not target_done): if not init_done: state = nixl_agent2.check_xfer_state(xfer_handle_2) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": init_done = True - print("Initiator done") + logger.info("Initiator done") if not target_done: if nixl_agent1.check_remote_xfer_done("initiator", b"UUID2"): target_done = True - print("Target done") + logger.info("Target done") nixl_agent2.release_xfer_handle(xfer_handle_1) nixl_agent2.release_xfer_handle(xfer_handle_2) @@ -205,4 +209,4 @@ nixl_utils.free_passthru(addr1) nixl_utils.free_passthru(addr3) - print("Test Complete.") + logger.info("Test Complete.") diff --git a/examples/python/nixl_gds_example.py b/examples/python/nixl_gds_example.py index c0d9c0c5c1..065968ca1e 100755 --- a/examples/python/nixl_gds_example.py +++ b/examples/python/nixl_gds_example.py @@ -20,17 +20,21 @@ import nixl._utils as nixl_utils from nixl._api import nixl_agent, nixl_agent_config +from nixl.logging import get_logger + +# Configure logging +logger = get_logger(__name__) + if __name__ == "__main__": buf_size = 16 * 4096 # Allocate memory and register with NIXL if len(sys.argv) < 2: - print("Please specify file path in argv") + logger.error("Please specify file path in argv") exit(0) - print("Using NIXL Plugins from:") - print(os.environ["NIXL_PLUGIN_DIR"]) + logger.info("Using NIXL Plugins from:\n%s", os.environ["NIXL_PLUGIN_DIR"]) agent_config = nixl_agent_config(backends=[]) nixl_agent1 = nixl_agent("GDSTester", agent_config) @@ -38,16 +42,19 @@ plugin_list = nixl_agent1.get_plugin_list() assert "GDS" in plugin_list - print("Plugin parameters") - print(nixl_agent1.get_plugin_mem_types("GDS")) - print(nixl_agent1.get_plugin_params("GDS")) + logger.info( + "Plugin parameters:\n%s\n%s\n", + nixl_agent1.get_plugin_mem_types("GDS"), + nixl_agent1.get_plugin_params("GDS"), + ) nixl_agent1.create_backend("GDS") - print("\nLoaded backend parameters") - print(nixl_agent1.get_backend_mem_types("GDS")) - print(nixl_agent1.get_backend_params("GDS")) - print() + logger.info( + "Backend parameters:\n%s\n%s\n", + nixl_agent1.get_backend_mem_types("GDS"), + nixl_agent1.get_backend_params("GDS"), + ) # get DRAM buf and initialize it to 0xba for verification addr1 = nixl_utils.malloc_passthru(buf_size) @@ -78,7 +85,7 @@ "WRITE", agent1_xfer1_descs, agent1_xfer_files, "GDSTester" ) if not xfer_handle_1: - print("Creating transfer failed.") + logger.error("Creating transfer failed.") exit() state = nixl_agent1.transfer(xfer_handle_1) @@ -89,18 +96,18 @@ while not done: state = nixl_agent1.check_xfer_state(xfer_handle_1) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": done = True - print("Initiator done") + logger.info("Initiator done") # read file data back into second buffer xfer_handle_2 = nixl_agent1.initialize_xfer( "READ", agent1_xfer2_descs, agent1_xfer_files, "GDSTester" ) if not xfer_handle_2: - print("Creating transfer failed.") + logger.error("Creating transfer failed.") exit() state = nixl_agent1.transfer(xfer_handle_2) @@ -111,11 +118,11 @@ while not done: state = nixl_agent1.check_xfer_state(xfer_handle_2) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": done = True - print("Initiator done") + logger.info("Initiator done") # transfer verification nixl_utils.verify_transfer(addr1, addr2, buf_size) @@ -130,4 +137,4 @@ os.close(agent1_fd) - print("Test Complete.") + logger.info("Test Complete.") diff --git a/examples/python/partial_md_example.py b/examples/python/partial_md_example.py index d84a1a762d..790a3fa411 100755 --- a/examples/python/partial_md_example.py +++ b/examples/python/partial_md_example.py @@ -21,6 +21,10 @@ import nixl._utils as nixl_utils from nixl._api import nixl_agent, nixl_agent_config from nixl._bindings import nixlNotFoundError +from nixl.logging import get_logger + +# Configure logging +logger = get_logger(__name__) def exchange_target_metadata( @@ -75,13 +79,14 @@ def invalidate_target_metadata( ) args = parser.parse_args() - print("Using NIXL Plugins from:") - print(os.environ["NIXL_PLUGIN_DIR"]) + logger.info("Using NIXL Plugins from: %s", os.environ["NIXL_PLUGIN_DIR"]) if args.etcd: etcd_endpoints = os.getenv("NIXL_ETCD_ENDPOINTS", "") if etcd_endpoints: - print("NIXL_ETCD_ENDPOINTS is set, using endpoints: ", etcd_endpoints) + logger.info( + "NIXL_ETCD_ENDPOINTS is set, using endpoints: %s", etcd_endpoints + ) else: raise ValueError( "NIXL_ETCD_ENDPOINTS is not set, but --etcd flag is provided" @@ -89,7 +94,7 @@ def invalidate_target_metadata( else: etcd_endpoints = "" del os.environ["NIXL_ETCD_ENDPOINTS"] - print("NIXL_ETCD_ENDPOINTS is not set, using socket exchange") + logger.info("NIXL_ETCD_ENDPOINTS is not set, using socket exchange") # Needed for socket exchange ip_addr = "127.0.0.1" @@ -114,8 +119,8 @@ def invalidate_target_metadata( target_strs2.append((addr1, buf_size, 0, "test")) malloc_addrs.append(addr1) - target_reg_descs1 = target_agent.get_reg_descs(target_strs1, "DRAM", is_sorted=True) - target_reg_descs2 = target_agent.get_reg_descs(target_strs2, "DRAM", is_sorted=True) + target_reg_descs1 = target_agent.get_reg_descs(target_strs1, "DRAM") + target_reg_descs2 = target_agent.get_reg_descs(target_strs2, "DRAM") target_xfer_descs1 = target_reg_descs1.trim() target_xfer_descs2 = target_reg_descs2.trim() @@ -132,7 +137,7 @@ def invalidate_target_metadata( init_strs.append((addr1, buf_size, 0, "test")) malloc_addrs.append(addr1) - init_reg_descs = init_agent.get_reg_descs(init_strs, "DRAM", is_sorted=True) + init_reg_descs = init_agent.get_reg_descs(init_strs, "DRAM") init_xfer_descs = init_reg_descs.trim() assert init_agent.register_memory(init_reg_descs) is not None @@ -167,16 +172,16 @@ def invalidate_target_metadata( if not init_done: state = init_agent.check_xfer_state(xfer_handle_1) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": init_done = True - print("Initiator done") + logger.info("Initiator done") if not target_done: if target_agent.check_remote_xfer_done("initiator", b"UUID1"): target_done = True - print("Target done") + logger.info("Target done") # Second set of descs was not sent, should fail try: @@ -184,9 +189,9 @@ def invalidate_target_metadata( "READ", init_xfer_descs, target_xfer_descs2, "target", b"UUID1" ) except nixlNotFoundError: - print("Correct exception") + logger.info("Correct exception") else: - print("Incorrect success") + logger.error("Incorrect success") os.abort() # Now send rest of descs @@ -207,7 +212,7 @@ def invalidate_target_metadata( try: # initialize transfer mode xfer_handle_2 = init_agent.initialize_xfer( - "READ", init_xfer_descs, target_xfer_descs1, "target", b"UUID1" + "READ", init_xfer_descs, target_xfer_descs2, "target", b"UUID1" ) except nixlNotFoundError: ready = False @@ -224,16 +229,16 @@ def invalidate_target_metadata( if not init_done: state = init_agent.check_xfer_state(xfer_handle_2) if state == "ERR": - print("Transfer got to Error state.") + logger.error("Transfer got to Error state.") exit() elif state == "DONE": init_done = True - print("Initiator done") + logger.info("Initiator done") if not target_done: if target_agent.check_remote_xfer_done("initiator", b"UUID1"): target_done = True - print("Target done") + logger.info("Target done") init_agent.release_xfer_handle(xfer_handle_1) init_agent.release_xfer_handle(xfer_handle_2) @@ -249,4 +254,4 @@ def invalidate_target_metadata( del init_agent del target_agent - print("Test Complete.") + logger.info("Test Complete.") diff --git a/examples/python/query_mem_example.py b/examples/python/query_mem_example.py index 9cf38cb84b..b108aaf999 100755 --- a/examples/python/query_mem_example.py +++ b/examples/python/query_mem_example.py @@ -21,18 +21,24 @@ try: from nixl._api import nixl_agent, nixl_agent_config + from nixl.logging import get_logger + + logger = get_logger(__name__) NIXL_AVAILABLE = True except ImportError: - print("NIXL API missing install NIXL.") + import logging + + logger = logging.getLogger(__name__) + logger.error("NIXL API missing install NIXL.") NIXL_AVAILABLE = False if __name__ == "__main__": - print("NIXL queryMem Python API Example") - print("=" * 40) + logger.info("NIXL queryMem Python API Example") + logger.info("=" * 40) if not NIXL_AVAILABLE: - print("Skipping example - NIXL bindings not available") + logger.warning("Skipping example - NIXL bindings not available") sys.exit(0) # Create temporary test files @@ -49,17 +55,18 @@ non_existent_file = "/tmp/nixl_example_nonexistent.txt" try: - print("Using NIXL Plugins from:") - print(os.environ["NIXL_PLUGIN_DIR"]) + logger.info("Using NIXL Plugins from: %s", os.environ["NIXL_PLUGIN_DIR"]) # Create an NIXL agent - print("Creating NIXL agent...") - config = nixl_agent_config(False, False, 0, []) + logger.info("Creating NIXL agent...") + config = nixl_agent_config( + enable_prog_thread=False, enable_listen_thread=False, backends=[] + ) agent = nixl_agent("example_agent", config) # Prepare a list of tuples as file paths in metaInfo field for querying. # Addr and length and devID fields are set to 0 for file queries. - print("Preparing file paths for querying...") + logger.info("Preparing file paths for querying...") file_paths = [ (0, 0, 0, temp_files[0]), # Existing file 1 (0, 0, 0, temp_files[1]), # Existing file 2 @@ -68,57 +75,57 @@ ] # Query memory using queryMem - print("Querying memory/storage information...") + logger.info("Querying memory/storage information...") # Try to create a backend with POSIX plugin try: params = agent.get_plugin_params("POSIX") agent.create_backend("POSIX", params) - print("Created backend: POSIX") + logger.info("Created backend: POSIX") # Query with specific backend resp = agent.query_memory(file_paths, "POSIX", mem_type="FILE") except Exception as e: - print(f"POSIX backend creation failed: {e}") + logger.exception("POSIX backend creation failed: %s", e) # Try MOCK_DRAM as fallback try: params = agent.get_plugin_params("MOCK_DRAM") agent.create_backend("MOCK_DRAM", params) - print("Created backend: MOCK_DRAM") + logger.info("Created backend: MOCK_DRAM") # Query with specific backend resp = agent.query_memory(file_paths, "MOCK_DRAM", mem_type="FILE") except Exception as e2: - print(f"MOCK_DRAM also failed: {e2}") - print("No working backends available") + logger.exception("MOCK_DRAM also failed: %s", e2) + logger.exception("No working backends available") sys.exit(0) # Display results - print(f"\nQuery results ({len(resp)} responses):") - print("-" * 50) + logger.info("\nQuery results (%d responses):", len(resp)) + logger.info("-" * 50) for i, result in enumerate(resp): - print(f"Descriptor {i}:") + logger.info("Descriptor %d:", i) if result is not None: - print(f" File size: {result.get('size', 'N/A')} bytes") - print(f" File mode: {result.get('mode', 'N/A')}") - print(f" Modified time: {result.get('mtime', 'N/A')}") + logger.info(" File size: %s bytes", result.get("size", "N/A")) + logger.info(" File mode: %s", result.get("mode", "N/A")) + logger.info(" Modified time: %s", result.get("mtime", "N/A")) else: - print(" File does not exist or is not accessible") - print() + logger.info(" File does not exist or is not accessible") + logger.info("") - print("Example completed successfully!") + logger.info("Example completed successfully!") except Exception as e: - print(f"Error in example: {e}") + logger.exception("Error in example: %s", e) import traceback traceback.print_exc() finally: # Clean up temporary files - print("Cleaning up temporary files...") + logger.info("Cleaning up temporary files...") for temp_file_path in temp_files: if os.path.exists(temp_file_path): os.unlink(temp_file_path) - print(f"Removed: {temp_file_path}") + logger.info("Removed: %s", temp_file_path) diff --git a/examples/python/telemetry_reader.py b/examples/python/telemetry_reader.py new file mode 100755 index 0000000000..e56598f1dd --- /dev/null +++ b/examples/python/telemetry_reader.py @@ -0,0 +1,303 @@ +#!/usr/bin/env python3 + +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import argparse +import ctypes +import logging +import mmap +import os +import signal +import sys +import time +from datetime import datetime + +# Set up logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(levelname)s - %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", +) +logger = logging.getLogger(__name__) + +# Constants from telemetry_event.h +TELEMETRY_VERSION = 1 +MAX_EVENT_NAME_LEN = 32 + +# NIXL telemetry categories +NIXL_TELEMETRY_MEMORY = 0 +NIXL_TELEMETRY_TRANSFER = 1 +NIXL_TELEMETRY_CONNECTION = 2 +NIXL_TELEMETRY_BACKEND = 3 +NIXL_TELEMETRY_ERROR = 4 +NIXL_TELEMETRY_PERFORMANCE = 5 +NIXL_TELEMETRY_SYSTEM = 6 +NIXL_TELEMETRY_CUSTOM = 7 +NIXL_TELEMETRY_MAX = 8 + +# Global flag for graceful shutdown +running = True + + +def signal_handler(signum, _): + """Signal handler for Ctrl+C""" + global running + if signum == signal.SIGINT: + logger.info("\nReceived Ctrl+C, shutting down...") + running = False + + +class NixlTelemetryEvent(ctypes.Structure): + """Python equivalent of nixlTelemetryEvent struct""" + + _pack_ = 1 + _fields_ = [ + ("timestamp_us", ctypes.c_uint64), + ("category", ctypes.c_int), + ("event_name", ctypes.c_char * MAX_EVENT_NAME_LEN), + ("_padding", ctypes.c_uint32), + ("value", ctypes.c_uint64), + ] + + +class BufferHeader(ctypes.Structure): + """Python equivalent of BufferHeader struct from cyclic_buffer.h""" + + _pack_ = 1 + _fields_ = [ + ("write_pos", ctypes.c_size_t), + ("read_pos", ctypes.c_size_t), + ("version", ctypes.c_uint32), + ("expected_version", ctypes.c_uint32), + ("capacity", ctypes.c_size_t), + ("mask", ctypes.c_size_t), + ] + + +class SharedRingBuffer: + """Python wrapper for the C++ SharedRingBuffer class""" + + def __init__(self, file_path, version=1): + self.file_path = file_path + self.version = version + self.file_fd = -1 + self.mmap_obj = None + self.header = None + self.data = None + self.buffer_size = None + + self._open_file() + self._map_memory() + + def _open_file(self): + """Open existing file""" + self.file_fd = os.open(self.file_path, os.O_RDWR) + + def _map_memory(self): + """Map the file into memory""" + self._map_header_only() + + def _map_header_only(self): + """Map only the header to read buffer size""" + # Map just the header first + header_mmap = mmap.mmap( + self.file_fd, + ctypes.sizeof(BufferHeader), + mmap.MAP_SHARED, + mmap.PROT_READ | mmap.PROT_WRITE, + ) + + temp_header = BufferHeader.from_buffer(header_mmap) + + if temp_header.version != self.version: + del temp_header + header_mmap.close() + raise RuntimeError( + f"Version mismatch: expected {self.version}, got {temp_header.version}" + ) + + self.buffer_size = temp_header.capacity + logger.info("Auto-detected buffer size: %d", self.buffer_size) + + del temp_header + header_mmap.close() + + # Now map the entire buffer + total_size = ( + ctypes.sizeof(BufferHeader) + + ctypes.sizeof(NixlTelemetryEvent) * self.buffer_size + ) + self.mmap_obj = mmap.mmap( + self.file_fd, total_size, mmap.MAP_SHARED, mmap.PROT_READ | mmap.PROT_WRITE + ) + + # Create ctypes pointers to the mapped memory + self.header = BufferHeader.from_buffer(self.mmap_obj) + data_offset = ctypes.sizeof(BufferHeader) + self.data = (NixlTelemetryEvent * self.buffer_size).from_buffer( + self.mmap_obj, data_offset + ) + + def get_version(self): + """Get the buffer version""" + return self.header.version + + def size(self): + """Get the number of events in the buffer""" + write_pos = self.header.write_pos + read_pos = self.header.read_pos + return (write_pos - read_pos) & self.header.mask + + def get_capacity(self): + """Get the buffer capacity""" + return self.buffer_size + + def empty(self): + """Check if buffer is empty""" + return self.header.read_pos == self.header.write_pos + + def full(self): + """Check if buffer is full""" + write_pos = self.header.write_pos + next_write = (write_pos + 1) & self.header.mask + return next_write == self.header.read_pos + + def pop(self): + """Pop an event from the buffer""" + read_pos = self.header.read_pos + + if read_pos == self.header.write_pos: + return None + + event = self.data[read_pos] + + next_read = (read_pos + 1) & self.header.mask + self.header.read_pos = next_read + + return event + + def __del__(self): + """Cleanup resources""" + # if self.mmap_obj: + # self.mmap_obj.close() + if self.file_fd != -1: + os.close(self.file_fd) + + +def format_timestamp(timestamp_us): + """Format timestamp in microseconds to readable format""" + dt = datetime.fromtimestamp(timestamp_us / 1_000_000) + microseconds = timestamp_us % 1_000_000 + return f"{dt.strftime('%Y-%m-%d %H:%M:%S')}.{microseconds:06d}" + + +def format_bytes(bytes_val): + """Format bytes to human readable format""" + units = ["B", "KB", "MB", "GB", "TB"] + unit_index = 0 + value = float(bytes_val) + + while value >= 1024.0 and unit_index < 4: + value /= 1024.0 + unit_index += 1 + + return f"{value:.2f} {units[unit_index]}" + + +def get_telemetry_category_string(category): + """Get string representation of telemetry category""" + category_strings = { + NIXL_TELEMETRY_MEMORY: "MEMORY", + NIXL_TELEMETRY_TRANSFER: "TRANSFER", + NIXL_TELEMETRY_CONNECTION: "CONNECTION", + NIXL_TELEMETRY_BACKEND: "BACKEND", + NIXL_TELEMETRY_ERROR: "ERROR", + NIXL_TELEMETRY_PERFORMANCE: "PERFORMANCE", + NIXL_TELEMETRY_SYSTEM: "SYSTEM", + NIXL_TELEMETRY_CUSTOM: "CUSTOM", + } + return category_strings.get(category, f"UNKNOWN_CATEGORY_{category}") + + +def print_telemetry_event(event): + """Print telemetry event in a formatted way""" + logger.info("\n=== NIXL Telemetry Event ===") + logger.info("Timestamp: %s", format_timestamp(event.timestamp_us)) + + # Decode event name + event_name = event.event_name.decode("utf-8").rstrip("\x00") + category_str = get_telemetry_category_string(event.category) + + logger.info("Category: %s", category_str) + logger.info("Event: %s", event_name) + logger.info("Value: %s", event.value) + logger.info("===========================") + + +def main(): + """Main function""" + parser = argparse.ArgumentParser(description="NIXL Telemetry Reader") + parser.add_argument( + "--telemetry_path", help="Path to the telemetry file", required=True + ) + + args = parser.parse_args() + + logger.info("Telemetry path: %s", args.telemetry_path) + telemetry_file_name = args.telemetry_path + if not os.path.exists(telemetry_file_name): + logger.error("Telemetry file %s does not exist", telemetry_file_name) + return 1 + + signal.signal(signal.SIGINT, signal_handler) + + try: + logger.info("Opening telemetry buffer: %s", telemetry_file_name) + logger.info("Press Ctrl+C to stop reading telemetry...") + + buffer = SharedRingBuffer(telemetry_file_name, version=TELEMETRY_VERSION) + + logger.info( + "Successfully opened telemetry buffer (version: %d)", buffer.get_version() + ) + logger.info("Buffer capacity: %d events", buffer.get_capacity()) + logger.info("Current events in buffer: %d", buffer.size()) + logger.info("Event structure size: %d bytes", ctypes.sizeof(NixlTelemetryEvent)) + + event_count = 0 + + while running: + # Try to read an event from the buffer + event = buffer.pop() + if event: + event_count += 1 + print_telemetry_event(event) + else: + # No events available, sleep briefly + time.sleep(0.1) + + logger.info("\nTotal events read: %d", event_count) + logger.info("Final buffer size: %d events", buffer.size()) + + except Exception as e: + logger.error("Error: %s", e) + return 1 + + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/examples/rust/Cargo.lock b/examples/rust/Cargo.lock index a04b44d8a3..19e594e248 100644 --- a/examples/rust/Cargo.lock +++ b/examples/rust/Cargo.lock @@ -432,7 +432,7 @@ dependencies = [ [[package]] name = "nixl-sys" -version = "0.5.0" +version = "0.6.0" dependencies = [ "bindgen", "cc", diff --git a/examples/rust/src/single_process_example.rs b/examples/rust/src/single_process_example.rs index 2639ac0fd5..cfed82bfe6 100644 --- a/examples/rust/src/single_process_example.rs +++ b/examples/rust/src/single_process_example.rs @@ -182,8 +182,7 @@ fn main() -> Result<(), Box> { if !completed { match agent1.get_xfer_status(&xfer_req) { Ok(status) => { - completed = !status; - if completed { + if status.is_success() { debug!("Transfer completed!"); } } diff --git a/meson.build b/meson.build index 6b27c9e4b0..adf891c6ed 100644 --- a/meson.build +++ b/meson.build @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -project('nixl', 'CPP', version: '0.5.0', +project('nixl', 'CPP', version: '0.6.0', default_options: ['buildtype=debug', 'werror=true', 'cpp_std=c++17', @@ -23,6 +23,7 @@ project('nixl', 'CPP', version: '0.5.0', # set up some global vars for compiler, platform, configuration, etc. cpp = meson.get_compiler('cpp') +fs = import('fs') dl_dep = cpp.find_library('dl', required: true) rt_dep = cpp.find_library('rt', required: true) @@ -52,6 +53,24 @@ cuda_stub_path = get_option('cudapath_stub') if cuda_lib_path == '' cuda_dep = dependency('cuda', required : false, modules : [ 'cudart', 'cuda' ]) + if not cuda_dep.found() + # Meson is not detecting CUDA reliably on ARM, fallback to default + cuda_home = run_command('bash', '-c', 'echo $CUDA_HOME').stdout().strip() + if cuda_home == '' + cuda_home = '/usr/local/cuda' + endif + cuda_lib = cuda_home + '/lib64' + cuda_inc = cuda_home + '/include' + cuda_stub = cuda_lib + '/stubs' + if fs.exists(cuda_lib) and fs.exists(cuda_inc) + cuda_dep = declare_dependency( + link_args : ['-L' + cuda_lib, '-L' + cuda_stub, '-lcuda', '-lcudart'], + include_directories : include_directories(cuda_inc)) + if cuda_dep.found() + message('Found CUDA installation through fallback method:', cuda_home) + endif + endif + endif else message('cuda lib path ', cuda_lib_path) if cuda_stub_path == '' @@ -81,6 +100,8 @@ if cuda_dep.found() nvcc_flags_link += ['-gencode=arch=compute_80,code=sm_80'] nvcc_flags_link += ['-gencode=arch=compute_90,code=sm_90'] add_project_link_arguments(nvcc_flags_link, language: 'cuda') +else + warning('CUDA not found. UCX backend will be built without CUDA support, and some plugins will be disabled.') endif # DOCA @@ -130,6 +151,45 @@ else ucx_dep = dependency('ucx', modules: ['ucx::ucs', 'ucx::ucp', 'ucx::uct']) endif +libfabric_path = get_option('libfabric_path') +if libfabric_path != '' + libfabric_lib_path = libfabric_path + '/lib' + libfabric_inc_path = libfabric_path + '/include' + libfabric_dep = declare_dependency( + link_args : ['-L' + libfabric_lib_path, '-lfabric'], + include_directories : include_directories(libfabric_inc_path)) +else + libfabric_dep = dependency('libfabric', required: false) +endif + +# UCX GPU device API detection +nvcc_prog = find_program('nvcc', required: false) +ucx_gpu_device_api_available = false +if ucx_dep.found() and cuda_dep.found() and nvcc_prog.found() + cuda = meson.get_compiler('cuda') + have_gpu_side = cuda.compiles(''' + #include + int main() { return 0; } + ''', dependencies : ucx_dep, args: nvcc_flags) + + have_host_side = cpp.compiles(''' + #include + int main() { return 0; } + ''', dependencies: ucx_dep) + + if have_gpu_side and have_host_side + ucx_gpu_device_api_available = true + add_project_arguments('-DHAVE_UCX_GPU_DEVICE_API', language: ['cpp', 'cuda']) + endif + + summary({ + 'UCX GPU Device API' : ucx_gpu_device_api_available, + 'GPU-side compile' : have_gpu_side, + 'Host-side compile' : have_host_side, + 'nvcc available' : nvcc_prog.found(), + }, section: 'UCX GPU Device API', bool_yn: true) +endif + if get_option('disable_gds_backend') add_project_arguments('-DDISABLE_GDS_BACKEND', language: 'cpp') endif @@ -163,6 +223,7 @@ if get_option('buildtype') == 'debug' endif nixl_inc_dirs = include_directories('src/api/cpp', 'src/api/cpp/backend', 'src/infra', 'src/core') +nixl_gpu_inc_dirs = include_directories('src/api/gpu/ucx') plugins_inc_dirs = include_directories('src/plugins') utils_inc_dirs = include_directories('src/utils') @@ -185,6 +246,9 @@ if get_option('install_headers') install_headers('src/core/transfer_request.h', install_dir: prefix_inc) install_headers('src/core/agent_data.h', install_dir: prefix_inc) install_headers('src/infra/mem_section.h', install_dir: prefix_inc) + if ucx_gpu_device_api_available + install_headers('src/api/gpu/ucx/nixl_device.cuh', install_dir: prefix_inc + '/gpu/ucx') + endif endif # Doxygen documentation diff --git a/meson_options.txt b/meson_options.txt index 3a5280a8f2..a316184f8d 100644 --- a/meson_options.txt +++ b/meson_options.txt @@ -14,6 +14,7 @@ # limitations under the License. option('ucx_path', type: 'string', value: '', description: 'Path to UCX install') +option('libfabric_path', type: 'string', value: '', description: 'Path to LIBFABRIC install') option('etcd_inc_path', type: 'string', value: '', description: 'Path to ETCD Headers') option('etcd_lib_path', type: 'string', value: '', description: 'Path to ETCD Libraries') option('disable_gds_backend', type : 'boolean', value : false, description : 'disable gds backend') @@ -26,6 +27,7 @@ option('cudapath_stub', type: 'string', value: '', description: 'Extra Stub path option('static_plugins', type: 'string', value: '', description: 'Plugins to be built-in, comma-separated') option('build_docs', type: 'boolean', value: false, description: 'Build Doxygen documentation') option('log_level', type: 'combo', choices: ['trace', 'debug', 'info', 'warning', 'error', 'fatal', 'auto'], value: 'auto', description: 'Log Level (auto: auto-detect based on build type: trace for debug builds, info for release builds)') +option('rust', type: 'boolean', value: false, description: 'Build Rust bindings') # Tests option('test_all_plugins', type: 'boolean', value: false, description: 'Testing all plugins in addition to the mocks..') diff --git a/nixl.pc.in b/nixl.pc.in new file mode 100644 index 0000000000..127a3daed9 --- /dev/null +++ b/nixl.pc.in @@ -0,0 +1,13 @@ +prefix=@prefix@ +exec_prefix=@prefix@ +libdir=@prefix@/@libdir@ +includedir=@prefix@/@includedir@ + +Name: nixl +Description: NVIDIA Inference Xfer Library +Version: @version@ +Requires: ucx +Requires.private: etcd-cpp-api +Libs: -L${libdir} -lnixl -lnixl_build -lnixl_common -lstream -lserdes -lucx_utils -lstdc++ +Libs.private: -ldl -lrt -lpthread +Cflags: -I${includedir} -I${includedir}/utils/serdes -I${includedir}/utils/common -I${includedir}/backend diff --git a/pyproject.toml b/pyproject.toml index 46854812c4..3cfbde6ef2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,7 +19,7 @@ build-backend = "mesonpy" [project] name = 'nixl' -version = '0.5.0' +version = '0.6.0' description = 'NIXL Python API' readme = 'README.md' license = {file = 'LICENSE'} diff --git a/src/api/cpp/backend/backend_aux.h b/src/api/cpp/backend/backend_aux.h index b6e2037f2b..aa2c6f216c 100644 --- a/src/api/cpp/backend/backend_aux.h +++ b/src/api/cpp/backend/backend_aux.h @@ -52,6 +52,7 @@ class nixlBackendInitParams { bool enableProgTh; nixlTime::us_t pthrDelay; nixl_thread_sync_t syncMode; + bool enableTelemetry_; }; // Pure virtual class to have a common pointer type diff --git a/src/api/cpp/backend/backend_engine.h b/src/api/cpp/backend/backend_engine.h index 7ea3a5ec81..44c852f0f8 100644 --- a/src/api/cpp/backend/backend_engine.h +++ b/src/api/cpp/backend/backend_engine.h @@ -20,8 +20,14 @@ #include #include #include +#include +#include + #include "nixl_types.h" #include "backend_aux.h" +#include "telemetry_event.h" + +constexpr size_t MAX_TELEMETRY_QUEUE_SIZE = 1000; // Base backend engine class for different backend implementations class nixlBackendEngine { @@ -29,11 +35,14 @@ class nixlBackendEngine { // Members that cannot be modified by a child backend and parent bookkeep nixl_backend_t backendType; nixl_b_params_t customParams; + std::vector telemetryEvents_; + std::mutex telemetryEventsMutex_; protected: // Members that can be accessed by the child (localAgent cannot be modified) bool initErr = false; const std::string localAgent; + const bool enableTelemetry_; [[nodiscard]] nixl_status_t setInitParam(const std::string &key, const std::string &value) { @@ -52,12 +61,25 @@ class nixlBackendEngine { return NIXL_ERR_INVALID_PARAM; } + void + addTelemetryEvent(const std::string &event_name, uint64_t value) { + if (!enableTelemetry_) return; + if (telemetryEvents_.size() >= MAX_TELEMETRY_QUEUE_SIZE) return; + std::lock_guard lock(telemetryEventsMutex_); + telemetryEvents_.emplace_back(std::chrono::duration_cast( + std::chrono::system_clock::now().time_since_epoch()) + .count(), + nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND, + event_name, + value); + } + public: - explicit nixlBackendEngine (const nixlBackendInitParams* init_params) + explicit nixlBackendEngine(const nixlBackendInitParams *init_params) : backendType(init_params->type), customParams(*init_params->customParams), - localAgent(init_params->localAgent) { - } + localAgent(init_params->localAgent), + enableTelemetry_(init_params->enableTelemetry_) {} nixlBackendEngine(nixlBackendEngine&&) = delete; nixlBackendEngine(const nixlBackendEngine&) = delete; @@ -67,6 +89,12 @@ class nixlBackendEngine { virtual ~nixlBackendEngine() = default; + std::vector + getTelemetryEvents() { + std::lock_guard lock(telemetryEventsMutex_); + return std::move(telemetryEvents_); + } + bool getInitErr() const noexcept { return initErr; } const nixl_backend_t& getType() const noexcept { return backendType; } const nixl_b_params_t& getCustomParams() const noexcept { return customParams; } @@ -84,9 +112,6 @@ class nixlBackendEngine { // pure virtual, and return errors, as parent shouldn't call if supportsNotif is false. virtual bool supportsNotif() const = 0; - // Determines if a backend supports progress thread. - virtual bool supportsProgTh() const = 0; - virtual nixl_mem_list_t getSupportedMems() const = 0; // TODO: Return by const-reference and mark noexcept? @@ -131,6 +156,30 @@ class nixlBackendEngine { //Backend aborts the transfer if necessary, and destructs the relevant objects virtual nixl_status_t releaseReqH(nixlBackendReqH* handle) const = 0; + // Create a GPU transfer request to GPU memory for GPU transfer. + virtual nixl_status_t + createGpuXferReq(const nixlBackendReqH &req_hndl, + const nixl_meta_dlist_t &local_descs, + const nixl_meta_dlist_t &remote_descs, + nixlGpuXferReqH &gpu_req_hndl) const { + return NIXL_ERR_NOT_SUPPORTED; + } + + // Release a GPU transfer request from GPU memory + virtual void + releaseGpuXferReq(nixlGpuXferReqH gpu_req_hndl) const {} + + // Get the size required for a GPU signal + virtual nixl_status_t + getGpuSignalSize(size_t &signal_size) const { + return NIXL_ERR_NOT_SUPPORTED; + } + + // Initialize a signal for GPU transfer using memory handle from descriptor + virtual nixl_status_t + prepGpuSignal(const nixlBackendMD &meta, void *signal) const { + return NIXL_ERR_NOT_SUPPORTED; + } // *** Needs to be implemented if supportsRemote() is true *** // @@ -181,14 +230,6 @@ class nixlBackendEngine { } - // *** Needs to be implemented if supportsProgTh() is true *** // - - // Force backend engine worker to progress. - virtual int - progress() { - return 0; - } - // *** Optional virtual methods that are good to be implemented in any backend *** // // Query information about a list of memory/storage diff --git a/src/api/cpp/backend/backend_plugin.h b/src/api/cpp/backend/backend_plugin.h index 68f289b552..15e17b05c3 100644 --- a/src/api/cpp/backend/backend_plugin.h +++ b/src/api/cpp/backend/backend_plugin.h @@ -19,6 +19,10 @@ #define __BACKEND_PLUGIN_H #include "backend/backend_engine.h" +#include "common/nixl_log.h" + +// Forward declarations for special engine types +class nixlUcxEngine; // Define the plugin API version #define NIXL_PLUGIN_API_VERSION 1 @@ -50,16 +54,71 @@ class nixlBackendPlugin { // Macro to define exported C functions for the plugin #define NIXL_PLUGIN_EXPORT __attribute__((visibility("default"))) +// Template for creating backend plugins with minimal boilerplate +template class nixlBackendPluginCreator { +public: + static nixlBackendPlugin * + create(int api_version, + const char *name, + const char *version, + const nixl_b_params_t ¶ms, + const nixl_mem_list_t &mem_list) { + + static const char *plugin_name = name; + static const char *plugin_version = version; + static const nixl_b_params_t plugin_params = params; + static const nixl_mem_list_t plugin_mems = mem_list; + + static nixlBackendPlugin plugin_instance = {api_version, + createEngine, + destroyEngine, + []() { return plugin_name; }, + []() { return plugin_version; }, + []() { return plugin_params; }, + []() { return plugin_mems; }}; + + return &plugin_instance; + } + +private: + [[nodiscard]] static nixlBackendEngine * + createEngine(const nixlBackendInitParams *init_params) { + try { + if constexpr (std::is_same_v) { + // UCX engine uses a factory pattern + auto engine = EngineType::create(*init_params); + return engine.release(); + } else { + // Other engines use direct constructor + return new EngineType(init_params); + } + } + catch (const std::exception &e) { + NIXL_ERROR << "Failed to create engine: " << e.what(); + return nullptr; + } + } + + static void + destroyEngine(nixlBackendEngine *engine) { + delete engine; + } +}; + + // Creator Function type for static plugins typedef nixlBackendPlugin* (*nixlStaticPluginCreatorFunc)(); // Plugin must implement these functions for dynamic loading +// Note: extern "C" is required for dynamic loading to avoid C++ name mangling extern "C" { - // Initialize the plugin - NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init(); +// Initialize the plugin +NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init(); - // Cleanup the plugin - NIXL_PLUGIN_EXPORT void nixl_plugin_fini(); +// Cleanup the plugin +NIXL_PLUGIN_EXPORT void +nixl_plugin_fini(); } #endif // __BACKEND_PLUGIN_H diff --git a/src/api/cpp/nixl.h b/src/api/cpp/nixl.h index f698601797..956ad4fe7d 100644 --- a/src/api/cpp/nixl.h +++ b/src/api/cpp/nixl.h @@ -170,8 +170,7 @@ class nixlAgent { * - For loopback descriptors, it is set to local agent's name, indicating that * this is for a loopback (local) transfer to be uued for remote_side handle * If a list of backends hints is provided (via extra_params), the preparation - * is limited to the specified backends. If `descs` has the sorted flag, that - * enables an optimization to speed up the preparation process. + * is limited to the specified backends. * * @param agent_name Agent name as a string for preparing xfer handle * @param descs The descriptor list to be prepared for transfer requests @@ -226,8 +225,6 @@ class nixlAgent { * pre-processing done in the preparation step. If a list of backends hints is * provided (via extra_params), the selection is limited to the specified backends. * Optionally, a notification message can also be provided through extra_params. - * If `local_descs` or `remote_descs` have the sorted flag, that enables an - * optimization to speed up the preparation process. * * @param operation Operation for transfer (e.g., NIXL_WRITE) * @param local_descs Local descriptor list @@ -289,6 +286,17 @@ class nixlAgent { nixl_status_t getXferStatus (nixlXferReqH* req_hndl) const; + + /** + * @brief Get the telemetry data associated with `req_hndl`. + * + * @param req_hndl Transfer request handle obtained from makeXferReq/createXferReq + * @param telemetry [out] Output telemetry information + * @return nixl_status_t Error code if call was not successful + */ + nixl_status_t + getXferTelemetry(const nixlXferReqH *req_hndl, nixl_xfer_telem_t &telemetry) const; + /** * @brief Query the backend associated with `req_hndl`. E.g., if for genNotif * the same backend as a transfer is desired. @@ -311,6 +319,55 @@ class nixlAgent { nixl_status_t releaseXferReq (nixlXferReqH* req_hndl) const; + /** + * @brief Create a GPU transfer request from a transfer request. + * + * @param req_hndl [in] Transfer request obtained from makeXferReq/createXferReq + * @param gpu_req_hndl [out] GPU transfer request handle + * @return nixl_status_t Error code if call was not successful + */ + nixl_status_t + createGpuXferReq(const nixlXferReqH &req_hndl, nixlGpuXferReqH &gpu_req_hndl) const; + + /** + * @brief Release transfer request from GPU memory + * + * @param gpu_req_hndl [in] GPU transfer request handle to be released + */ + void + releaseGpuXferReq(nixlGpuXferReqH gpu_req_hndl) const; + + /** + * @brief Get the size required for a GPU signal. + * + * This function returns the size required for allocating memory for a GPU signal. + * The returned size should be used to allocate memory that will be registered + * and used with @ref prepGpuSignal. + * + * @param signal_size [out] Size required for the GPU signal + * @param extra_params Extra parameters used in getting the size of the GPU signal. + * The backend must be specified in extra_params. + * @return nixl_status_t Error code if call was not successful + */ + nixl_status_t + getGpuSignalSize(size_t &signal_size, const nixl_opt_args_t *extra_params) const; + + /** + * @brief Prepare a signal for GPU transfer. + * + * The caller must allocate and register the signal memory before calling this function. + * Use @ref getGpuSignalSize to query the required signal size, allocate + * the signal accordingly, and register it using @ref registerMem. + * + * @param signal_descs [in] Registered descriptor list for the signal memory + * @param extra_params Extra parameters used in preparing the GPU signal. + * The backend must be specified in extra_params. + * @return nixl_status_t Error code if call was not successful + */ + nixl_status_t + prepGpuSignal(const nixl_reg_dlist_t &signal_descs, + const nixl_opt_args_t *extra_params) const; + /** * @brief Release the prepared descriptor list handle `dlist_hndl` * diff --git a/src/api/cpp/nixl_descriptors.h b/src/api/cpp/nixl_descriptors.h index 4eb072c7f0..8ebb3af91f 100644 --- a/src/api/cpp/nixl_descriptors.h +++ b/src/api/cpp/nixl_descriptors.h @@ -30,96 +30,102 @@ * element, alongside supporting methods */ class nixlBasicDesc { - public: - /** @var Start of Buffer */ - uintptr_t addr; - /** @var Buffer Length */ - size_t len; - /** @var deviceID/blockID/fileID */ - uint64_t devId; +public: + /** @var Start of Buffer */ + uintptr_t addr; + /** @var Buffer Length */ + size_t len; + /** @var deviceID/blockID/fileID */ + uint64_t devId; - /** - * @brief Default constructor for nixlBasicDesc - * Does not initialize members to zero - */ - nixlBasicDesc() {}; - /** - * @brief Parametrized constructor for nixlBasicDesc - * - * @param addr Start of buffer/block/offset-in-file - * @param len Length of buffer - * @param devID deviceID/BlockID/fileID - */ - nixlBasicDesc(const uintptr_t &addr, - const size_t &len, - const uint64_t &dev_id); - /** - * @brief Deserializer constructor for nixlBasicDesc with - * serialized blob of another nixlBasicDesc - * - * @param str Serialized Descriptor - */ - nixlBasicDesc(const nixl_blob_t &str); // deserializer - /** - * @brief Copy constructor for nixlBasicDesc - * - * @param desc nixlBasicDesc object - */ - nixlBasicDesc(const nixlBasicDesc &desc) = default; - /** - * @brief Operator (=) overloading constructor - * with nixlBasicDesc object - * - * @param desc nixlBasicDesc object - */ - nixlBasicDesc& operator=(const nixlBasicDesc &desc) = default; - /** - * @brief nixlBasicDesc destructor - */ - ~nixlBasicDesc() = default; - /** - * @brief Operator overloading (<) to compare BasicDesc objects - * Comparison criteria is devID, then addr, then len - */ - bool operator<(const nixlBasicDesc &desc) const; - /** - * @brief Operator overloading (==) to compare BasicDesc objects - * - * @param lhs nixlBasicDesc object - * @param rhs nixlBasicDesc object - * - */ - friend bool operator==(const nixlBasicDesc &lhs, const nixlBasicDesc &rhs); - /** - * @brief Operator overloading (!=) to compare BasicDesc objects - * - * @param lhs nixlBasicDesc object - * @param rhs nixlBasicDesc object - * - */ - friend bool operator!=(const nixlBasicDesc &lhs, const nixlBasicDesc &rhs); - /** - * @brief Check if current object address range covers the input object's - * - * @param query nixlBasicDesc object - */ - bool covers (const nixlBasicDesc &query) const; - /** - * @brief Check for overlap between BasicDesc objects - * - * @param query nixlBasicDesc Object - */ - bool overlaps (const nixlBasicDesc &query) const; - /** - * @brief Serialize descriptor into a blob - */ - nixl_blob_t serialize() const; - /** - * @brief Print descriptor for debugging - * - * @param suffix gets prepended to the descriptor print - */ - void print(const std::string &suffix) const; + /** + * @brief Default constructor for nixlBasicDesc + * Does not initialize members to zero + */ + nixlBasicDesc() {}; + /** + * @brief Parametrized constructor for nixlBasicDesc + * + * @param addr Start of buffer/block/offset-in-file + * @param len Length of buffer + * @param devID deviceID/BlockID/fileID + */ + nixlBasicDesc(const uintptr_t &addr, const size_t &len, const uint64_t &dev_id); + /** + * @brief Deserializer constructor for nixlBasicDesc with + * serialized blob of another nixlBasicDesc + * + * @param str Serialized Descriptor + */ + nixlBasicDesc(const nixl_blob_t &str); // deserializer + /** + * @brief Copy constructor for nixlBasicDesc + * + * @param desc nixlBasicDesc object + */ + nixlBasicDesc(const nixlBasicDesc &desc) = default; + /** + * @brief Operator (=) overloading constructor + * with nixlBasicDesc object + * + * @param desc nixlBasicDesc object + */ + nixlBasicDesc & + operator=(const nixlBasicDesc &desc) = default; + /** + * @brief nixlBasicDesc destructor + */ + ~nixlBasicDesc() = default; + /** + * @brief Operator overloading (<) to compare BasicDesc objects + * Comparison criteria is devID, then addr, then len + */ + bool + operator<(const nixlBasicDesc &desc) const; + /** + * @brief Operator overloading (==) to compare BasicDesc objects + * + * @param lhs nixlBasicDesc object + * @param rhs nixlBasicDesc object + * + */ + friend bool + operator==(const nixlBasicDesc &lhs, const nixlBasicDesc &rhs); + /** + * @brief Operator overloading (!=) to compare BasicDesc objects + * + * @param lhs nixlBasicDesc object + * @param rhs nixlBasicDesc object + * + */ + friend bool + operator!=(const nixlBasicDesc &lhs, const nixlBasicDesc &rhs); + /** + * @brief Check if current object address range covers the input object's + * + * @param query nixlBasicDesc object + */ + bool + covers(const nixlBasicDesc &query) const; + /** + * @brief Check for overlap between BasicDesc objects + * + * @param query nixlBasicDesc Object + */ + bool + overlaps(const nixlBasicDesc &query) const; + /** + * @brief Serialize descriptor into a blob + */ + nixl_blob_t + serialize() const; + /** + * @brief Print descriptor for debugging + * + * @param suffix gets prepended to the descriptor print + */ + void + print(const std::string &suffix) const; }; /** @@ -128,54 +134,58 @@ class nixlBasicDesc { * bundled with a nixlBasicDesc. */ class nixlBlobDesc : public nixlBasicDesc { - public: - /** @var blob for metadata information */ - nixl_blob_t metaInfo; +public: + /** @var blob for metadata information */ + nixl_blob_t metaInfo; - /** @var Reuse parent constructor without the metadata */ - using nixlBasicDesc::nixlBasicDesc; + /** @var Reuse parent constructor without the metadata */ + using nixlBasicDesc::nixlBasicDesc; - /** - * @brief Parametrized constructor for nixlBlobDesc - * - * @param addr Start of buffer/block/offset-in-file - * @param len Length of buffer - * @param devID deviceID/BlockID/fileID - * @param meta_info Metadata blob - */ - nixlBlobDesc(const uintptr_t &addr, const size_t &len, - const uint64_t &dev_id, const nixl_blob_t &meta_info); - /** - * @brief Constructor for nixlBlobDesc from nixlBasicDesc and metadata blob - * - * @param desc nixlBasicDesc object - * @param meta_info Metadata blob - */ - nixlBlobDesc(const nixlBasicDesc &desc, const nixl_blob_t &meta_info); - /** - * @brief Deserializer constructor for nixlBlobDesc with serialized blob - * - * @param str Serialized blob from another nixlBlobDesc - */ - nixlBlobDesc(const nixl_blob_t &str); - /** - * @brief Operator overloading (==) to compare nixlBlobDesc objects - * - * @param lhs nixlBlobDesc object - * @param rhs nixlBlobDesc object - */ - friend bool operator==(const nixlBlobDesc &lhs, - const nixlBlobDesc &rhs); - /** - * @brief Serialize nixlBlobDesc to a blob - */ - nixl_blob_t serialize() const; - /** - * @brief Print nixlBlobDesc for debugging purpose - * - * @param suffix gets prepended to the descriptor print - */ - void print(const std::string &suffix) const; + /** + * @brief Parametrized constructor for nixlBlobDesc + * + * @param addr Start of buffer/block/offset-in-file + * @param len Length of buffer + * @param devID deviceID/BlockID/fileID + * @param meta_info Metadata blob + */ + nixlBlobDesc(const uintptr_t &addr, + const size_t &len, + const uint64_t &dev_id, + const nixl_blob_t &meta_info); + /** + * @brief Constructor for nixlBlobDesc from nixlBasicDesc and metadata blob + * + * @param desc nixlBasicDesc object + * @param meta_info Metadata blob + */ + nixlBlobDesc(const nixlBasicDesc &desc, const nixl_blob_t &meta_info); + /** + * @brief Deserializer constructor for nixlBlobDesc with serialized blob + * + * @param str Serialized blob from another nixlBlobDesc + */ + nixlBlobDesc(const nixl_blob_t &str); + /** + * @brief Operator overloading (==) to compare nixlBlobDesc objects + * + * @param lhs nixlBlobDesc object + * @param rhs nixlBlobDesc object + */ + friend bool + operator==(const nixlBlobDesc &lhs, const nixlBlobDesc &rhs); + /** + * @brief Serialize nixlBlobDesc to a blob + */ + nixl_blob_t + serialize() const; + /** + * @brief Print nixlBlobDesc for debugging purpose + * + * @param suffix gets prepended to the descriptor print + */ + void + print(const std::string &suffix) const; }; /** @@ -185,154 +195,175 @@ class nixlBlobDesc : public nixlBasicDesc { */ template class nixlDescList { - private: - /** @var NIXL memory type */ - nixl_mem_t type; - /** @var Flag for if list should be kept sorted - * Comparison is done based on nixlBasicDesc (<) operator which - * has comparison order of devID, then addr, then len. - */ - bool sorted; - /** @var Vector for storing nixlDescs */ - std::vector descs; +protected: + /** @var NIXL memory type */ + nixl_mem_t type; + /** @var Vector for storing nixlDescs */ + std::vector descs; - public: - /** - * @brief Parametrized Constructor for nixlDescList - * - * @param type NIXL memory type of descriptor list - * @param sorted Flag to set sorted option (default = false) - * @param init_size initial size for descriptor list (default = 0) - */ - nixlDescList(const nixl_mem_t &type, - const bool &sorted=false, - const int &init_size=0); - /** - * @brief Deserializer constructor for nixlDescList from nixlSerDes object - * which serializes/deserializes our classes into/from blobs - * - * @param deserialize nixlSerDes object to construct nixlDescList - */ - nixlDescList(nixlSerDes* deserializer); - /** - * @brief Copy constructor for creating nixlDescList from another object - * of the same type. - * - * @param d_list other nixlDescList object of the same type - */ - nixlDescList(const nixlDescList &d_list) = default; - /** - * @brief Operator = overloading constructor for nixlDescList - * - * @param d_list nixlDescList object - */ - nixlDescList& operator=(const nixlDescList &d_list) = default; - /** - * @brief nixlDescList Destructor - */ - ~nixlDescList () = default; - /** - * @brief Get NIXL memory type for this DescList - */ - inline nixl_mem_t getType() const { return type; } - /** - * @brief get sorted flag - */ - inline bool isSorted() const { return sorted; } - /** - * @brief Get count of descriptors - */ - inline int descCount() const { return descs.size(); } - /** - * @brief Check if nixlDescList is empty or not - */ - inline bool isEmpty() const { return (descs.size()==0); } - /** - * @brief Check if any two nixlDescs in the internal list of descriptors - * overlap with each other - */ - bool hasOverlaps() const; - /** - * @brief Operator [] overloading, get/set descriptor at [index]. - * Can throw std::out_of_range exception. - */ - const T& operator[](unsigned int index) const; - T& operator[](unsigned int index); - /** - * @brief Vector like iterators for const and non-const elements - */ - inline typename std::vector::const_iterator begin() const - { return descs.begin(); } - inline typename std::vector::const_iterator end() const - { return descs.end(); } - inline typename std::vector::iterator begin() - { return descs.begin(); } - inline typename std::vector::iterator end() - { return descs.end(); } - /** - * @brief Operator overloading (==) to compare nixlDescList objects - * - * @param lhs nixlDescList object - * @param rhs nixlDescList object - * - */ - template friend bool operator==(const nixlDescList &lhs, - const nixlDescList &rhs); - /** - * @brief Resize nixlDescList object. If new size is more than the - * original size, the sorted status will be negated if set. - * - * @param count Number of elements after resizing DescList object - */ - void resize (const size_t &count); - /** - * @brief Verify if a nixlDescList is sorted, for instance after using - * resize and adding new elements. If true, the sorted flag is set. - */ - bool verifySorted(); - /** - * @brief Empty the descriptors list - */ - inline void clear() { descs.clear(); } - /** - * @brief Add Descriptors to descriptor list - * If nixlDescList object is sorted, this method keeps it sorted - */ - void addDesc(const T &desc); - /** - * @brief Remove descriptor from list at index - * Can throw std::out_of_range exception. - */ - void remDesc(const int &index); - /** - * @brief Convert a nixlDescList with metadata by trimming it to a - * nixlDescList of nixlBasicDesc elements - */ - nixlDescList trim() const; - /** - * @brief Check if input descriptor `desc` overlaps with any descriptor - * within the current object, and returns its index if found. - * - * @param index [out] index of overlapping descriptor - */ - bool overlaps (const T &desc, int &index) const; - /** - * @brief Get the index of a descriptor that matches the `query` - * - * @param query nixlBasicDesc object to find among the object's descriptors - * @return int index of the queried nixlBasicDesc if found, or negative error value - */ - int getIndex(const nixlBasicDesc &query) const; - /** - * @brief Serialize a descriptor list with nixlSerDes class - * @param serializer nixlSerDes object to serialize nixlDescList - * @return nixl_status_t Error code if serialize was not successful - */ - nixl_status_t serialize(nixlSerDes* serializer) const; - /** - * @brief Print the descriptor list for debugging - */ - void print() const; +public: + /** + * @brief Parametrized Constructor for nixlDescList + * + * @param type NIXL memory type of descriptor list + * @param init_size initial size for descriptor list (default = 0) + */ + nixlDescList(const nixl_mem_t &type, const int &init_size = 0); + + /** + * @brief Deserializer constructor for nixlDescList from nixlSerDes object + * which serializes/deserializes our classes into/from blobs + * + * @param deserialize nixlSerDes object to construct nixlDescList + */ + nixlDescList(nixlSerDes *deserializer); + + /** + * @brief Copy constructor for creating nixlDescList from another object + * of the same type. + * + * @param d_list other nixlDescList object of the same type + */ + nixlDescList(const nixlDescList &d_list) = default; + + /** + * @brief Operator = overloading constructor for nixlDescList + * + * @param d_list nixlDescList object + */ + nixlDescList & + operator=(const nixlDescList &d_list) = default; + + /** + * @brief nixlDescList Destructor + */ + virtual ~nixlDescList() = default; + + /** + * @brief Get NIXL memory type for this DescList + */ + inline nixl_mem_t + getType() const { + return type; + } + + /** + * @brief Get count of descriptors + */ + inline int + descCount() const { + return descs.size(); + } + + /** + * @brief Check if nixlDescList is empty or not + */ + inline bool + isEmpty() const { + return (descs.size() == 0); + } + + /** + * @brief Operator [] overloading, get/set descriptor at [index]. + * Can throw std::out_of_range exception. + */ + const T & + operator[](unsigned int index) const; + virtual T & + operator[](unsigned int index); + + /** + * @brief Vector like iterators for const and non-const elements + */ + inline typename std::vector::const_iterator + begin() const { + return descs.begin(); + } + + inline typename std::vector::const_iterator + end() const { + return descs.end(); + } + + inline typename std::vector::iterator + begin() { + return descs.begin(); + } + + inline typename std::vector::iterator + end() { + return descs.end(); + } + + /** + * @brief Operator overloading (==) to compare nixlDescList objects + * + * @param lhs nixlDescList object + * @param rhs nixlDescList object + * + */ + template + friend bool + operator==(const nixlDescList &lhs, const nixlDescList &rhs); + + /** + * @brief Resize nixlDescList object. + * + * @param count Number of elements after resizing DescList object + */ + virtual void + resize(const size_t &count); + + /** + * @brief Empty the descriptors list + */ + inline void + clear() { + descs.clear(); + } + + /** + * @brief Add Descriptors to descriptor list + */ + virtual void + addDesc(const T &desc); + + /** + * @brief Remove descriptor from list at index + * Can throw std::out_of_range exception. + */ + void + remDesc(const int &index); + + /** + * @brief Convert a nixlDescList with metadata by trimming it to a + * nixlDescList of nixlBasicDesc elements + */ + nixlDescList + trim() const; + + /** + * @brief Get the index of a descriptor that matches the `query` + * + * @param query nixlBasicDesc object to find among the object's descriptors + * @return int index of the queried nixlBasicDesc if found, or negative error value + */ + virtual int + getIndex(const nixlBasicDesc &query) const; + + /** + * @brief Serialize a descriptor list with nixlSerDes class + * @param serializer nixlSerDes object to serialize nixlDescList + * @return nixl_status_t Error code if serialize was not successful + */ + nixl_status_t + serialize(nixlSerDes *serializer) const; + + /** + * @brief Print the descriptor list for debugging + */ + void + print() const; }; /** * @brief A typedef for a nixlDescList diff --git a/src/api/cpp/nixl_params.h b/src/api/cpp/nixl_params.h index 7655ee4702..d5869b6bb3 100644 --- a/src/api/cpp/nixl_params.h +++ b/src/api/cpp/nixl_params.h @@ -37,6 +37,8 @@ class nixlAgentConfig { int listenPort; /** @var synchronization mode for multi-threaded environment execution */ nixl_thread_sync_t syncMode; + /** @var Capture telemetry info regardless of environment variables*/ + bool captureTelemetry; public: @@ -53,29 +55,42 @@ class nixlAgentConfig { */ uint64_t lthrDelay; + /** + * @var ETCD watch timeout in microseconds + * Timeout for waiting for metadata changes when watching etcd keys. + */ + std::chrono::microseconds etcdWatchTimeout; /** * @brief Agent configuration constructor for enabling various features. * @param use_prog_thread flag to determine use of progress thread - * @param use_listen_thread flag to determine use of listener thread - * @param port specify port for listener thread to listen on - * @param sync_mode Thread synchronization mode + * @param use_listen_thread Optional flag to determine use of listener thread + * @param port Optional port for listener thread to listen on + * @param sync_mode Optional Thread synchronization mode + * @param num_workers Optional number of shared workers per backend * @param pthr_delay_us Optional delay for pthread in us * @param lthr_delay_us Optional delay for listener thread in us + * @param capture_telemetry Optional flag to enable telemetry capture + * @param etcd_watch_timeout Optional timeout for etcd watch operations in microseconds */ - nixlAgentConfig (const bool use_prog_thread, - const bool use_listen_thread=false, - const int port=0, - nixl_thread_sync_t sync_mode=nixl_thread_sync_t::NIXL_THREAD_SYNC_DEFAULT, - unsigned int num_workers = 1, - const uint64_t pthr_delay_us=0, - const uint64_t lthr_delay_us = 100000) : - useProgThread(use_prog_thread), - useListenThread(use_listen_thread), - listenPort(port), - syncMode(sync_mode), - pthrDelay(pthr_delay_us), - lthrDelay(lthr_delay_us) { } + nixlAgentConfig(const bool use_prog_thread, + const bool use_listen_thread = false, + const int port = 0, + nixl_thread_sync_t sync_mode = nixl_thread_sync_t::NIXL_THREAD_SYNC_DEFAULT, + unsigned int num_workers = 1, + const uint64_t pthr_delay_us = 0, + const uint64_t lthr_delay_us = 100000, + const bool capture_telemetry = false, + const std::chrono::microseconds &etcd_watch_timeout = + std::chrono::microseconds(5000000)) + : useProgThread(use_prog_thread), + useListenThread(use_listen_thread), + listenPort(port), + syncMode(sync_mode), + captureTelemetry(capture_telemetry), + pthrDelay(pthr_delay_us), + lthrDelay(lthr_delay_us), + etcdWatchTimeout(etcd_watch_timeout) {} /** * @brief Copy constructor for nixlAgentConfig object diff --git a/src/api/cpp/nixl_types.h b/src/api/cpp/nixl_types.h index ef4475b1ef..30a9d4a348 100644 --- a/src/api/cpp/nixl_types.h +++ b/src/api/cpp/nixl_types.h @@ -20,6 +20,7 @@ #include #include #include +#include /*** Forward declarations ***/ @@ -61,7 +62,9 @@ enum nixl_status_t { NIXL_ERR_REPOST_ACTIVE = -7, NIXL_ERR_UNKNOWN = -8, NIXL_ERR_NOT_SUPPORTED = -9, - NIXL_ERR_REMOTE_DISCONNECT = -10 + NIXL_ERR_REMOTE_DISCONNECT = -10, + NIXL_ERR_CANCELED = -11, + NIXL_ERR_NO_TELEMETRY = -12 }; /** @@ -222,6 +225,60 @@ struct nixlAgentOptionalArgs { */ using nixl_opt_args_t = nixlAgentOptionalArgs; +/** + * @brief A typedef for a nixlGpuXferReqH + */ +using nixlGpuXferReqH = void *; + +/** + * @brief A typedefs for a point in time + */ +using chrono_point_t = std::chrono::steady_clock::time_point; + +/** + * @brief A typedefs for a period of time in microseconds + */ +using chrono_period_us_t = std::chrono::microseconds; + +/** + * @struct nixlXferTelemetry + * @brief A structure for telemetry output from agent API + */ +struct nixlXferTelemetry { + /** + * @var startTime Time that the transfer was posted + */ + chrono_point_t startTime; + + /** + * @var postDuration Time it took to do the post operation + */ + chrono_period_us_t postDuration; + + /** + * @var xferDuration Time it took to complete the transfer + * if checkXferReq is called late, that might impact this result + */ + chrono_period_us_t xferDuration; + + /** + * @var totalBytes Amount of bytes transferred in the request + */ + size_t totalBytes; + + /** + * @var descCount Number of descriptors in the transfer request. + * If any merging of descriptors were performed, it will be reflected here. + */ + size_t descCount; +}; + +/** + * @brief A typedef for a nixlXferTelemetry + * for telemetry output. + */ +using nixl_xfer_telem_t = nixlXferTelemetry; + /** * @brief A define for an empty string, that indicates the descriptor list is being * prepared for the local agent as an initiator in prepXferDlist method. diff --git a/src/api/gpu/ucx/nixl_device.cuh b/src/api/gpu/ucx/nixl_device.cuh new file mode 100644 index 0000000000..4c3edb9a02 --- /dev/null +++ b/src/api/gpu/ucx/nixl_device.cuh @@ -0,0 +1,272 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef _NIXL_DEVICE_CUH +#define _NIXL_DEVICE_CUH + +#include +#include + +struct nixlGpuXferStatusH { + ucp_device_request_t device_request; +}; + +struct nixlGpuSignal { + uint64_t inc = 0; + uint64_t remote_addr = 0; +}; + +/** + * @enum nixl_gpu_level_t + * @brief An enumeration of different levels for GPU transfer requests. + */ +enum class nixl_gpu_level_t : uint64_t { + THREAD = UCS_DEVICE_LEVEL_THREAD, + WARP = UCS_DEVICE_LEVEL_WARP, + BLOCK = UCS_DEVICE_LEVEL_BLOCK, + GRID = UCS_DEVICE_LEVEL_GRID +}; + +/** + * @brief Parameters for GPU transfer requests with safe type conversion. + */ +struct nixlGpuXferReqParams { + nixlGpuXferReqParams() = delete; + + __device__ + nixlGpuXferReqParams(nixlGpuXferReqH req_hndl, + bool is_no_delay, + nixlGpuXferStatusH *xfer_status) + : mem_list{static_cast(req_hndl)}, + flags{is_no_delay ? static_cast(UCP_DEVICE_FLAG_NODELAY) : 0}, + ucp_request{xfer_status ? &xfer_status->device_request : nullptr} {} + + ucp_device_mem_list_handle_h mem_list; + uint64_t flags; + ucp_device_request_t *ucp_request; +}; + +/** + * @brief Convert UCS status to NIXL status. + * + * @param status [in] UCS status code. + * + * @return nixl_status_t Corresponding NIXL status code. + */ +__device__ inline nixl_status_t +nixlGpuConvertUcsStatus(ucs_status_t status) { + return status == UCS_OK ? NIXL_SUCCESS : NIXL_ERR_BACKEND; +} + +/** + * @brief Post a memory transfer request to the GPU. + * + * @param req_hndl [in] Request handle. + * @param index [in] Index of the memory descriptor in the transfer request. + * @param addr [in] Local address of the memory to be transferred. + * @param remote_addr [in] Remote address of the memory to be transferred to. + * @param size [in] Size of the memory to be transferred. + * @param is_no_delay [in] Whether to use no-delay mode. + * @param xfer_status [out] Status of the transfer. If null, the status is not reported. + * + * @return nixl_status_t Error code if call was not successful + */ +template +__device__ nixl_status_t +nixlGpuPostSingleWriteXferReq(nixlGpuXferReqH req_hndl, + unsigned index, + const void *addr, + uint64_t remote_addr, + size_t size, + bool is_no_delay = true, + nixlGpuXferStatusH *xfer_status = nullptr) { + const nixlGpuXferReqParams params{req_hndl, is_no_delay, xfer_status}; + + ucs_status_t status = ucp_device_put_single(level)>( + params.mem_list, index, addr, remote_addr, size, params.flags, params.ucp_request); + + return nixlGpuConvertUcsStatus(status); +} + +/** + * @brief Post a signal transfer request to the GPU. + * + * @param req_hndl [in] Request handle. + * @param index [in] Index of the signal to be transferred. + * @param signal [in] Signal to be sent. + * @param is_no_delay [in] Whether to use no-delay mode. + * @param xfer_status [out] Status of the transfer. If null, the status is not reported. + * + * @return nixl_status_t Error code if call was not successful + */ +template +__device__ nixl_status_t +nixlGpuPostSignalXferReq(nixlGpuXferReqH req_hndl, + unsigned index, + const nixlGpuSignal &signal, + bool is_no_delay = true, + nixlGpuXferStatusH *xfer_status = nullptr) { + const nixlGpuXferReqParams params{req_hndl, is_no_delay, xfer_status}; + + ucs_status_t status = ucp_device_counter_inc(level)>( + params.mem_list, index, signal.inc, signal.remote_addr, params.flags, params.ucp_request); + + return nixlGpuConvertUcsStatus(status); +} + +/** + * @brief Post a partial memory transfer request to the GPU. + * + * @param req_hndl [in] Request handle. + * @param count [in] Number of blocks to send. This is also the length of the arrays + * @a indices, @a sizes, @a addrs, and @a remote_addrs. + * @param indices [in] Indices of the blocks to send. + * @param sizes [in] Sizes of the blocks to send. + * @param addrs [in] Addresses of the blocks to send. + * @param remote_addrs [in] Remote addresses of the blocks to send to. + * @param signal [in] Signal to be sent. + * @param is_no_delay [in] Whether to use no-delay mode. + * @param xfer_status [out] Status of the transfer. If null, the status is not reported. + * + * @return nixl_status_t Error code if call was not successful + */ +template +__device__ nixl_status_t +nixlGpuPostPartialWriteXferReq(nixlGpuXferReqH req_hndl, + size_t count, + const unsigned *indices, + const size_t *sizes, + void *const *addrs, + const uint64_t *remote_addrs, + const nixlGpuSignal &signal, + unsigned signal_index, + bool is_no_delay = true, + nixlGpuXferStatusH *xfer_status = nullptr) { + const nixlGpuXferReqParams params{req_hndl, is_no_delay, xfer_status}; + + ucs_status_t status = + ucp_device_put_multi_partial(level)>(params.mem_list, + indices, + count, + addrs, + remote_addrs, + sizes, + signal_index, + signal.inc, + signal.remote_addr, + params.flags, + params.ucp_request); + + return nixlGpuConvertUcsStatus(status); +} + +/** + * @brief Post a memory transfer request to the GPU. + * + * @param req_hndl [in] Request handle. + * @param sizes [in] Sizes of the blocks to send. + * @param addrs [in] Addresses of the blocks to send. + * @param remote_addrs [in] Remote addresses of the blocks to send to. + * @param signal [in] Signal to be sent. + * @param is_no_delay [in] Whether to use no-delay mode. + * @param xfer_status [out] Status of the transfer. If null, the status is not reported. + * + * @note The arrays @a sizes, @a addrs, and @a remote_addrs must have the same length, which + * corresponds to the number of blocks to transfer as specified in @a req_hndl. + * + * @return nixl_status_t Error code if call was not successful + */ +template +__device__ nixl_status_t +nixlGpuPostWriteXferReq(nixlGpuXferReqH req_hndl, + const size_t *sizes, + void *const *addrs, + const uint64_t *remote_addrs, + const nixlGpuSignal &signal, + bool is_no_delay = true, + nixlGpuXferStatusH *xfer_status = nullptr) { + const nixlGpuXferReqParams params{req_hndl, is_no_delay, xfer_status}; + + ucs_status_t status = + ucp_device_put_multi(level)>(params.mem_list, + addrs, + remote_addrs, + sizes, + signal.inc, + signal.remote_addr, + params.flags, + params.ucp_request); + + return nixlGpuConvertUcsStatus(status); +} + +/** + * @brief Get the status of the transfer request. + * + * @param xfer_status [in] Status of the transfer. + * + * @return NIXL_SUCCESS The request has completed, no more operations are in progress. + * @return NIXL_IN_PROG One or more operations in the request have not completed. + * @return Error code if call was not successful + */ +template +__device__ nixl_status_t +nixlGpuGetXferStatus(nixlGpuXferStatusH &xfer_status) { + const auto status = ucp_device_progress_req(level)>( + &xfer_status.device_request); + + switch (status) { + case UCS_OK: + return NIXL_SUCCESS; + case UCS_INPROGRESS: + return NIXL_IN_PROG; + default: + return NIXL_ERR_BACKEND; + } +} + +/** + * @brief Read the signal. + * + * The signal must be initialized with the host function @ref prepGpuSignal. + * + * @param signal [in] Address of the signal. + * + * @return The signal. + */ +template +__device__ uint64_t +nixlGpuReadSignal(const void *signal) { + return ucp_device_counter_read(level)>(signal); +} + +/** + * @brief Write value to the local signal. + * + * This function can be used to set a signal to a specific value. + * + * The signal must be initialized with the host function @ref prepGpuSignal. + * + * @param signal [in,out] Address of the signal. + * @param value [in] Value to write to the signal. + */ +template +__device__ void +nixlGpuWriteSignal(void *signal, uint64_t value) { + ucp_device_counter_write(level)>(signal, value); +} + +#endif // _NIXL_DEVICE_CUH diff --git a/src/api/python/_api.py b/src/api/python/_api.py index 58db3f2215..8923cbf1b6 100644 --- a/src/api/python/_api.py +++ b/src/api/python/_api.py @@ -20,13 +20,103 @@ import torch import nixl._bindings as nixlBind +from nixl.logging import get_logger + +# Get logger using centralized configuration +logger = get_logger(__name__) DEFAULT_COMM_PORT = nixlBind.DEFAULT_COMM_PORT -# Opaque nixl handle types + +""" +@brief Opaque handle wrapper for a prepared transfer descriptor list. + Use release() to explicitly free resources; __del__ performs best-effort cleanup. +@param agent Owning nixl_agent used to perform release operations. +@param value Internal handle +""" + + +class nixl_prepped_dlist_handle: + __slots__ = ("_handle", "_agent", "_released") + + def __init__(self, agent, value: int): + self._handle = int(value) + self._agent = agent + self._released = False + + def __repr__(self) -> str: + return ( + f"nixl_prepped_dlist_handle(0x{self._handle:x}, released={self._released})" + ) + + def release(self): + if not self._released: + self._agent.releasedDlistH(self._handle) + self._released = True + + def __del__(self): + if not self._released: + try: + self._agent.releasedDlistH(self._handle) + except Exception: + try: + logger.error( + "nixl_prepped_dlist_handle finalization failed for 0x%x", + self._handle, + ) + except Exception: + pass + + +""" +@brief Opaque handle wrapper for a transfer request. + Use release() to explicitly free resources. If transfer was not complete, this will initiate + the abort process (if available) and will raise an exception. + __del__ calls release() and if it fails, it logs the failure and defers release by queuing + the handle in leaked xfer handles list, which will be re-released during agent destruction +@param agent Owning nixl_agent used to perform release operations. +@param value Internal handle +""" + + +class nixl_xfer_handle: + __slots__ = ("_handle", "_agent", "_released") + + def __init__(self, agent, value: int): + self._handle = int(value) + self._agent = agent + self._released = False + + def __repr__(self) -> str: + return f"nixl_xfer_handle(0x{self._handle:x}, released={self._released})" + + def release(self): + if not self._released: + self._agent.releaseXferReq(self._handle) + self._released = True + + def __del__(self): + if not self._released: + try: + self._agent.releaseXferReq(self._handle) + except Exception: + try: + logger.error( + "nixl_xfer_handle finalization failed for 0x%x; keeping handle alive in agent leak list", + self._handle, + ) + except Exception: + pass + try: + self._agent._leaked_xfer_handles.append(self._handle) + except Exception: + pass + return + + +# Opaque handle for backend can be just int, as it's not passed to the user nixl_backend_handle = int -nixl_prepped_dlist_handle = int -nixl_xfer_handle = int + """ @brief Configuration class for NIXL agent. @@ -34,7 +124,8 @@ @param enable_prog_thread Whether to enable the progress thread, if available. @param enable_listen_thread Whether to enable the listener thread for metadata communication. @param listen_port Specify the port for the listener thread to listen on. - +@param capture_telemetry Whether to enable telemetry capture. +@param num_threads Specify number of threads for the supported multi-threaded backends. @param backends List of backend names for agent to initialize. Default is UCX, other backends can be added to the list, or after agent creation, can be initialized with create_backend. @@ -47,6 +138,8 @@ def __init__( enable_prog_thread: bool = True, enable_listen_thread: bool = False, listen_port: int = 0, + capture_telemetry: bool = False, + num_threads: int = 0, backends: list[str] = ["UCX"], ): # TODO: add backend init parameters @@ -54,6 +147,8 @@ def __init__( self.enable_pthread = enable_prog_thread self.enable_listen = enable_listen_thread self.port = listen_port + self.capture_telemetry = capture_telemetry + self.num_threads = num_threads """ @@ -76,7 +171,7 @@ def __init__( ): if nixl_conf and instantiate_all: instantiate_all = False - print( + logger.warning( "Ignoring instantiate_all based on the provided config in agent creation." ) if not nixl_conf: @@ -94,10 +189,15 @@ def __init__( nixl_conf.enable_listen, nixl_conf.port, thread_config, + 1, + 0, + 100000, + nixl_conf.capture_telemetry, ) self.agent = nixlBind.nixlAgent(agent_name, agent_config) self.name = agent_name + self._leaked_xfer_handles: list[int] = [] self.notifs: dict[str, list[bytes]] = {} self.backends: dict[str, nixl_backend_handle] = {} self.backend_mems: dict[str, list[str]] = {} @@ -105,7 +205,7 @@ def __init__( self.plugin_list = self.agent.getAvailPlugins() if len(self.plugin_list) == 0: - print("No plugins available, cannot start transfers!") + logger.error("No plugins available, cannot start transfers!") raise RuntimeError("No plugins available for NIXL, cannot start transfers!") self.plugin_b_options: dict[str, dict[str, str]] = {} @@ -115,23 +215,24 @@ def __init__( self.plugin_b_options[plugin] = backend_options self.plugin_mem_types[plugin] = mem_types - # TODO: populate init from default parameters, or define a set of params in python - init: dict[str, str] = {} - if instantiate_all: - for plugin in self.plugin_list: - self.create_backend(plugin, init) - else: - for bknd in nixl_conf.backends: - # TODO: populate init from nixl_conf when added - if bknd not in self.plugin_list: - print( - "Skipping backend registration", - bknd, - "due to the missing plugin.", - ) - else: - self.create_backend(bknd, init) + nixl_conf.backends = self.plugin_list + + for bknd in nixl_conf.backends: + if bknd not in self.plugin_list: + logger.warning( + "Skipping backend registration %s due to the missing plugin.", + bknd, + ) + else: + # TODO: improve population of init from nixl_conf + init: dict[str, str] = {} + if nixl_conf.num_threads > 0: + if bknd == "UCX" or bknd == "OBJ": + init["num_threads"] = str(nixl_conf.num_threads) + elif bknd == "GDS_MT": + init["thread_count"] = str(nixl_conf.num_threads) + self.create_backend(bknd, init) self.nixl_mems = { "DRAM": nixlBind.DRAM_SEG, @@ -147,7 +248,22 @@ def __init__( "READ": nixlBind.NIXL_READ, } - print("Initialized NIXL agent:", agent_name) + logger.info("Initialized NIXL agent: %s", agent_name) + + def __del__(self): + # Best-effort cleanup of any leaked xfer handles belonging to this agent + if getattr(self, "_leaked_xfer_handles", None): + for h in list(self._leaked_xfer_handles): + try: + self.releaseXferReq(h) + except Exception as e: + try: + logger.error( + "Failed to finalize leaked nixl_xfer_handle 0x%x: %s", h, e + ) + except Exception: + pass + self._leaked_xfer_handles.clear() """ @brief Get the list of available plugins. @@ -169,7 +285,9 @@ def get_plugin_mem_types(self, backend: str) -> list[str]: if backend in self.plugin_mem_types: return self.plugin_mem_types[backend] else: - print("Plugin", backend, "is not available to get its supported mem types.") + logger.warning( + "Plugin %s is not available to get its supported mem types.", backend + ) return [] """ @@ -184,7 +302,7 @@ def get_plugin_params(self, backend: str) -> dict[str, str]: if backend in self.plugin_b_options: return self.plugin_b_options[backend] else: - print("Plugin", backend, "is not available to get its parameters.") + logger.warning("Plugin %s is not available to get its parameters.", backend) return {} """ @@ -201,8 +319,8 @@ def get_backend_mem_types(self, backend: str) -> list[str]: if backend in self.backend_mems: return self.backend_mems[backend] else: - print( - "Backend", backend, "not instantiated to get its supported mem types." + logger.warning( + "Backend %s not instantiated to get its supported mem types.", backend ) return [] @@ -220,7 +338,9 @@ def get_backend_params(self, backend: str) -> dict[str, str]: if backend in self.backend_options: return self.backend_options[backend] else: - print("Backend", backend, "not instantiated to get its parameters.") + logger.warning( + "Backend %s not instantiated to get its parameters.", backend + ) return {} """ @@ -238,14 +358,13 @@ def create_backend(self, backend: str, initParams: dict[str, str] = {}): ) self.backend_mems[backend] = mem_types self.backend_options[backend] = backend_options - print("Backend", backend, "was instantiated") + logger.info("Backend %s was instantiated", backend) """ @brief Register memory regions, optionally with specified backends. @param reg_list List of either memory regions, tensors, or nixlRegDList to register. @param mem_type Optional memory type, necessary if specifying a list of memory regions. - @param is_sorted Optional bool for a list of memory regions or tensors for if they are sorted. @param backends Optional list of backend names for registration, otherwise NIXL will try to register with all backends that support this memory type. @return nixlRegDList for the registered memory, can be used with deregister_memory. @@ -255,10 +374,9 @@ def register_memory( self, reg_list, mem_type: Optional[str] = None, - is_sorted: bool = False, backends: list[str] = [], ) -> nixlBind.nixlRegDList: - reg_descs = self.get_reg_descs(reg_list, mem_type, is_sorted) + reg_descs = self.get_reg_descs(reg_list, mem_type) handle_list = [] for backend_string in backends: @@ -295,7 +413,7 @@ def deregister_memory( def query_memory( self, reg_list, backend: str, mem_type: Optional[str] = None ) -> list[Optional[dict[str, str]]]: - reg_descs = self.get_reg_descs(reg_list, mem_type, False) + reg_descs = self.get_reg_descs(reg_list, mem_type) # Get the backend handle if backend not in self.backends: @@ -339,8 +457,6 @@ def make_connection(self, remote_agent: str, backends: list[str] = []): @param xfer_list List of transfer descriptors, can be list of memory region tuples, tensors, Nx3 numpy array, or nixlXferDList. See get_xfer_descs for more details on the structure. @param mem_type Optional memory type necessary for list of memory regions. - @param is_sorted Optional bool for whether memory region list or tensor list are sorted. - For long lists of transfer descriptors, sorting can speed up transfer preparation. @param backends Optional list of backend names to limit which backends are used during preparation @return Opaque handle to the prepared transfer descriptor list. """ @@ -350,10 +466,9 @@ def prep_xfer_dlist( agent_name: str, xfer_list, mem_type: Optional[str] = None, - is_sorted: bool = False, backends: list[str] = [], ) -> nixl_prepped_dlist_handle: - descs = self.get_xfer_descs(xfer_list, mem_type, is_sorted) + descs = self.get_xfer_descs(xfer_list, mem_type) if agent_name == "NIXL_INIT_AGENT" or agent_name == "": agent_name = nixlBind.NIXL_INIT_AGENT @@ -363,8 +478,7 @@ def prep_xfer_dlist( handle_list.append(self.backends[backend_string]) handle = self.agent.prepXferDlist(agent_name, descs, handle_list) - - return handle + return nixl_prepped_dlist_handle(self.agent, handle) """ @brief Estimate the cost of a transfer operation. @@ -375,7 +489,7 @@ def prep_xfer_dlist( """ def estimate_xfer_cost(self, req_handle: nixl_xfer_handle) -> tuple[int, int, int]: - duration, err_margin, method = self.agent.estimateXferCost(req_handle) + duration, err_margin, method = self.agent.estimateXferCost(req_handle._handle) if method == nixlBind.NIXL_COST_ANALYTICAL_BACKEND: method = "ANALYTICAL_BACKEND" else: @@ -397,6 +511,7 @@ def estimate_xfer_cost(self, req_handle: nixl_xfer_handle) -> tuple[int, int, in @param backends Optional list of backend names to limit which backends NIXL can use. @param skip_desc_merge Whether to skip descriptor merging optimization. @return Opaque handle for posting/checking transfer. + The handle can be released by calling release_xfer_handle from agent, or release() method on itself. """ def make_prepped_xfer( @@ -417,16 +532,16 @@ def make_prepped_xfer( handle = self.agent.makeXferReq( op, - local_xfer_side, + local_xfer_side._handle, local_indices, - remote_xfer_side, + remote_xfer_side._handle, remote_indices, notif_msg, handle_list, skip_desc_merge, ) - return handle + return nixl_xfer_handle(self.agent, handle) """ @brief Initialize a transfer operation. This is a combined API, to create a transfer request @@ -443,6 +558,7 @@ def make_prepped_xfer( notif_msg should be bytes, as that is what will be returned to the target, but will work with str too. @param backends Optional list of backend names to limit which backends NIXL can use. @return Opaque handle for posting/checking transfer. + The handle can be released by calling release_xfer_handle from agent, or release() method on itself. """ def initialize_xfer( @@ -463,7 +579,7 @@ def initialize_xfer( op, local_descs, remote_descs, remote_agent, notif_msg, handle_list ) - return handle + return nixl_xfer_handle(self.agent, handle) """ @brief Initiate a data transfer operation. @@ -478,7 +594,7 @@ def initialize_xfer( """ def transfer(self, handle: nixl_xfer_handle, notif_msg: bytes = b"") -> str: - status = self.agent.postXferReq(handle, notif_msg) + status = self.agent.postXferReq(handle._handle, notif_msg) if status == nixlBind.NIXL_SUCCESS: return "DONE" elif status == nixlBind.NIXL_IN_PROG: @@ -494,7 +610,7 @@ def transfer(self, handle: nixl_xfer_handle, notif_msg: bytes = b"") -> str: """ def check_xfer_state(self, handle: nixl_xfer_handle) -> str: - status = self.agent.getXferStatus(handle) + status = self.agent.getXferStatus(handle._handle) if status == nixlBind.NIXL_SUCCESS: return "DONE" elif status == nixlBind.NIXL_IN_PROG: @@ -502,6 +618,22 @@ def check_xfer_state(self, handle: nixl_xfer_handle) -> str: else: return "ERR" + """ + @brief Get telemetry information of a transfer request. + The output object has three time values fields in microseconds + (startTime, postDuration, xferDuration), as well as integer totalBytes transferred + for the request, and integer descCount representing number of descriptors involved + (for example if there was some merging of descriptors). + + @param handle Handle to the transfer operation, from make_prepped_xfer or initialize_xfer. + @return nixlXferTelemetry object + """ + + def get_xfer_telemetry( + self, handle: nixl_xfer_handle + ) -> nixlBind.nixlXferTelemetry: + return self.agent.getXferTelemetry(handle._handle) + """ @brief Query the backend that was chosen for a transfer operation. @@ -510,7 +642,7 @@ def check_xfer_state(self, handle: nixl_xfer_handle) -> str: """ def query_xfer_backend(self, handle: nixl_xfer_handle) -> str: - b_handle = self.agent.queryXferBackend(handle) + b_handle = self.agent.queryXferBackend(handle._handle) # this works because there should not be multiple matching handles in the Dict return next( backendS @@ -527,7 +659,7 @@ def query_xfer_backend(self, handle: nixl_xfer_handle) -> str: """ def release_xfer_handle(self, handle: nixl_xfer_handle): - self.agent.releaseXferReq(handle) + handle.release() """ @brief Release a descriptor list handle, which internally frees the memory used for the handle. @@ -536,7 +668,7 @@ def release_xfer_handle(self, handle: nixl_xfer_handle): """ def release_dlist_handle(self, handle: nixl_prepped_dlist_handle): - self.agent.releasedDlistH(handle) + handle.release() """ @brief Get new notifications that have come to the agent. @@ -777,8 +909,6 @@ def check_remote_metadata( @param descs List of any of the above types @param mem_type Optional memory type necessary for (a). - @param is_sorted Optional bool for if the descriptors are sorted for (a) and (c) - sort criteria has the comparison order of devID, then addr, then len. @return Transfer descriptor list, nixlXferDList. """ @@ -786,36 +916,31 @@ def get_xfer_descs( self, descs, mem_type: Optional[str] = None, - is_sorted: bool = False, ) -> nixlBind.nixlXferDList: # can add check for DLPack input if isinstance(descs, nixlBind.nixlXferDList): return descs elif isinstance(descs, nixlBind.nixlRegDList): - print("RegList type detected for transfer, please use XferList") + logger.error("RegList type detected for transfer, please use XferList") new_descs = None elif isinstance(descs[0], tuple): if mem_type is not None and len(descs[0]) == 3: - new_descs = nixlBind.nixlXferDList( - self.nixl_mems[mem_type], descs, is_sorted - ) + new_descs = nixlBind.nixlXferDList(self.nixl_mems[mem_type], descs) elif mem_type is None: - print("Please specify a mem type if not using Tensors") + logger.error("Please specify a mem type if not using Tensors") new_descs = None else: - print("3-tuple list needed for transfer") + logger.error("3-tuple list needed for transfer") new_descs = None elif isinstance(descs, np.ndarray): if mem_type is not None and descs.ndim == 2 and descs.shape[1] == 3: - new_descs = nixlBind.nixlXferDList( - self.nixl_mems[mem_type], descs, is_sorted - ) + new_descs = nixlBind.nixlXferDList(self.nixl_mems[mem_type], descs) elif mem_type is None: - print("Please specify a mem type if not using Tensors") + logger.error("Please specify a mem type if not using Tensors") new_descs = None else: - print( + logger.error( "Nx3 shape required for transfer descriptor list from numpy array" ) new_descs = None @@ -828,12 +953,10 @@ def get_xfer_descs( if gpu_id == -1: # DRAM gpu_id = 0 new_descs = nixlBind.nixlXferDList( - self.nixl_mems[mem_type], - [(base_addr, region_len, gpu_id)], - is_sorted, + self.nixl_mems[mem_type], [(base_addr, region_len, gpu_id)] ) else: - print("Please use a list of contiguous Tensors") + logger.error("Please use a list of contiguous Tensors") new_descs = None elif isinstance(descs[0], torch.Tensor): # List[torch.Tensor]: tensor_type = descs[0].device @@ -843,7 +966,7 @@ def get_xfer_descs( if descs[i].device != tensor_type: return None if not descs[i].is_contiguous(): - print("Please use a list of contiguous Tensors") + logger.error("Please use a list of contiguous Tensors") return None base_addr = descs[i].data_ptr() region_len = descs[i].numel() * descs[i].element_size() @@ -852,9 +975,7 @@ def get_xfer_descs( gpu_id = 0 dlist[i, :] = (base_addr, region_len, gpu_id) mem_type = "cuda" if str(tensor_type).startswith("cuda") else "cpu" - new_descs = nixlBind.nixlXferDList( - self.nixl_mems[mem_type], dlist, is_sorted - ) + new_descs = nixlBind.nixlXferDList(self.nixl_mems[mem_type], dlist) else: new_descs = None @@ -871,8 +992,6 @@ def get_xfer_descs( @param descs List of any of the above types @param mem_type Optional memory type necessary for (a). - @param is_sorted Optional bool for if the descriptors are sorted for (a) and (c) - sort criteria has the comparison order of devID, then addr, then len. @return Registration descriptor list, nixlRegDList. """ @@ -880,36 +999,31 @@ def get_reg_descs( self, descs, mem_type: Optional[str] = None, - is_sorted: bool = False, ) -> nixlBind.nixlRegDList: # can add check for DLPack input if isinstance(descs, nixlBind.nixlRegDList): return descs elif isinstance(descs, nixlBind.nixlXferDList): - print("XferList type detected for registration, please use RegList") + logger.error("XferList type detected for registration, please use RegList") new_descs = None elif isinstance(descs[0], tuple): if mem_type is not None and len(descs[0]) == 4: - new_descs = nixlBind.nixlRegDList( - self.nixl_mems[mem_type], descs, is_sorted - ) + new_descs = nixlBind.nixlRegDList(self.nixl_mems[mem_type], descs) elif mem_type is None: - print("Please specify a mem type if not using Tensors") + logger.error("Please specify a mem type if not using Tensors") new_descs = None else: - print("4-tuple list needed for registration") + logger.error("4-tuple list needed for registration") new_descs = None elif isinstance(descs, np.ndarray): if mem_type is not None and descs.ndim == 2 and descs.shape[1] == 3: - new_descs = nixlBind.nixlRegDList( - self.nixl_mems[mem_type], descs, is_sorted - ) + new_descs = nixlBind.nixlRegDList(self.nixl_mems[mem_type], descs) elif mem_type is None: - print("Please specify a mem type if not using Tensors") + logger.error("Please specify a mem type if not using Tensors") new_descs = None else: - print( + logger.error( "Nx3 shape required for transfer descriptor list from numpy array" ) new_descs = None @@ -922,12 +1036,10 @@ def get_reg_descs( if gpu_id == -1: # DRAM gpu_id = 0 new_descs = nixlBind.nixlRegDList( - self.nixl_mems[mem_type], - [(base_addr, region_len, gpu_id, "")], - is_sorted, + self.nixl_mems[mem_type], [(base_addr, region_len, gpu_id, "")] ) else: - print("Please use a list of contiguous Tensors") + logger.error("Please use a list of contiguous Tensors") new_descs = None elif isinstance(descs[0], torch.Tensor): # List[torch.Tensor]: tensor_type = descs[0].device @@ -937,7 +1049,7 @@ def get_reg_descs( if descs[i].device != tensor_type: return None if not descs[i].is_contiguous(): - print("Please use a list of contiguous Tensors") + logger.error("Please use a list of contiguous Tensors") return None base_addr = descs[i].data_ptr() region_len = descs[i].numel() * descs[i].element_size() @@ -946,9 +1058,7 @@ def get_reg_descs( gpu_id = 0 dlist[i, :] = (base_addr, region_len, gpu_id) mem_type = "cuda" if str(tensor_type).startswith("cuda") else "cpu" - new_descs = nixlBind.nixlRegDList( - self.nixl_mems[mem_type], dlist, is_sorted - ) + new_descs = nixlBind.nixlRegDList(self.nixl_mems[mem_type], dlist) else: new_descs = None diff --git a/src/api/python/logging.py b/src/api/python/logging.py new file mode 100644 index 0000000000..d92e76df54 --- /dev/null +++ b/src/api/python/logging.py @@ -0,0 +1,113 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Centralized logging configuration for NIXL. + +Usage: + from nixl_logging import get_logger + + logger = get_logger(__name__) + logger.info("This is a log message") + +Or for backward compatibility: + import nixl_logging + import logging + + logger = logging.getLogger(__name__) + logger.info("This is a log message") +""" + +import logging +import logging.config +import os + +_logging_configured = False + +LOGGING_CONFIG = { + "version": 1, + "disable_existing_loggers": False, + "formatters": { + "simpleFormatter": { + "format": "%(asctime)s NIXL %(levelname)-7s %(filename)s:%(lineno)d %(message)s", + "datefmt": "%Y-%m-%d %H:%M:%S", + } + }, + "handlers": { + "consoleHandler": { + "class": "logging.StreamHandler", + "level": "INFO", + "formatter": "simpleFormatter", + "stream": "ext://sys.stdout", + } + }, + "loggers": { + "nixl": {"level": "INFO", "handlers": ["consoleHandler"], "propagate": False} + }, +} + + +def set_log_level_by_env(nixl_logger: logging.Logger) -> None: + # Override log level from environment variable if set + env_log_level = os.getenv("NIXL_LOG_LEVEL") + if env_log_level: + try: + # Convert string to logging level + numeric_level = getattr(logging, env_log_level.upper(), None) + if not isinstance(numeric_level, int): + raise ValueError(f"Invalid log level: {env_log_level}") + + # Set the level for the nixl logger + nixl_logger.setLevel(numeric_level) + + nixl_logger.info( + "Log level set to %s from NIXL_LOG_LEVEL environment variable", + env_log_level.upper(), + ) + except (ValueError, AttributeError) as e: + nixl_logger.warning( + "Invalid NIXL_LOG_LEVEL value '%s': %s. Using configuration default.", + env_log_level, + e, + ) + + +def setup_logging() -> None: + global _logging_configured + + if _logging_configured: + return + + # Use dictionary-based configuration + logging.config.dictConfig(LOGGING_CONFIG) + + nixl_logger = logging.getLogger("nixl") + set_log_level_by_env(nixl_logger) + + logging.raiseExceptions = os.getenv("NIXL_DEBUG_LOGGING", "").lower() in ( + "true", + "1", + "yes", + ) + + _logging_configured = True + + +def get_logger(name: str) -> logging.Logger: + setup_logging() + # Convert module name to nixl hierarchy + # e.g., '_api' -> 'nixl.api', 'test_nixl_bindings' -> 'nixl.test_nixl_bindings' + clean_name = name.lstrip("_") # Remove leading underscores + return logging.getLogger(f"nixl.{clean_name}") diff --git a/src/api/python/meson.build b/src/api/python/meson.build index 6f7eb80eda..828c814ef9 100644 --- a/src/api/python/meson.build +++ b/src/api/python/meson.build @@ -18,3 +18,4 @@ py = import('python').find_installation('python3', pure: false) py.install_sources('_api.py', subdir: ('nixl')) py.install_sources('__init__.py', subdir: ('nixl')) py.install_sources('py.typed', subdir: ('nixl')) +py.install_sources('logging.py', subdir: ('nixl')) diff --git a/src/bindings/meson.build b/src/bindings/meson.build index ff83a632d4..7234b0e5e6 100644 --- a/src/bindings/meson.build +++ b/src/bindings/meson.build @@ -14,4 +14,6 @@ # limitations under the License. subdir('python') -subdir('rust') +if get_option('rust') + subdir('rust') +endif diff --git a/src/bindings/python/nixl_bindings.cpp b/src/bindings/python/nixl_bindings.cpp index 6e3ca37cff..0442aee662 100644 --- a/src/bindings/python/nixl_bindings.cpp +++ b/src/bindings/python/nixl_bindings.cpp @@ -18,6 +18,7 @@ #include #include #include +#include #include #include @@ -64,14 +65,29 @@ class nixlRepostActiveError : public std::runtime_error { nixlRepostActiveError(const char *what) : runtime_error(what) {} }; +class nixlUnknownError : public std::runtime_error { +public: + nixlUnknownError(const char *what) : runtime_error(what) {} +}; + class nixlNotSupportedError : public std::runtime_error { public: nixlNotSupportedError(const char *what) : runtime_error(what) {} }; -class nixlUnknownError : public std::runtime_error { +class nixlRemoteDisconnectError : public std::runtime_error { public: - nixlUnknownError(const char *what) : runtime_error(what) {} + nixlRemoteDisconnectError(const char *what) : runtime_error(what) {} +}; + +class nixlCancelledError : public std::runtime_error { +public: + nixlCancelledError(const char *what) : runtime_error(what) {} +}; + +class nixlNoTelemetryError : public std::runtime_error { +public: + nixlNoTelemetryError(const char *what) : runtime_error(what) {} }; void @@ -108,6 +124,15 @@ throw_nixl_exception(const nixl_status_t &status) { case NIXL_ERR_NOT_SUPPORTED: throw nixlNotSupportedError(nixlEnumStrings::statusStr(status).c_str()); break; + case NIXL_ERR_REMOTE_DISCONNECT: + throw nixlRemoteDisconnectError(nixlEnumStrings::statusStr(status).c_str()); + break; + case NIXL_ERR_CANCELED: + throw nixlCancelledError(nixlEnumStrings::statusStr(status).c_str()); + break; + case NIXL_ERR_NO_TELEMETRY: + throw nixlNoTelemetryError(nixlEnumStrings::statusStr(status).c_str()); + break; default: throw std::runtime_error("BAD_STATUS"); } @@ -161,6 +186,22 @@ PYBIND11_MODULE(_bindings, m) { .value("NIXL_ERR_NOT_SUPPORTED", NIXL_ERR_NOT_SUPPORTED) .export_values(); + py::class_(m, "nixlXferTelemetry") + .def(py::init<>()) + .def_property_readonly("startTime", + [](const nixl_xfer_telem_t &t) { + return std::chrono::duration_cast( + t.startTime.time_since_epoch()) + .count(); + }) + .def_property_readonly("postDuration", + [](const nixl_xfer_telem_t &t) { return t.postDuration.count(); }) + .def_property_readonly("xferDuration", + [](const nixl_xfer_telem_t &t) { return t.xferDuration.count(); }) + .def_readonly("totalBytes", &nixl_xfer_telem_t::totalBytes) + .def_readonly("descCount", &nixl_xfer_telem_t::descCount); + + py::register_exception(m, "nixlNotPostedError"); py::register_exception(m, "nixlInvalidParamError"); py::register_exception(m, "nixlBackendError"); @@ -170,13 +211,13 @@ PYBIND11_MODULE(_bindings, m) { py::register_exception(m, "nixlRepostActiveError"); py::register_exception(m, "nixlUnknownError"); py::register_exception(m, "nixlNotSupportedError"); + py::register_exception(m, "nixlRemoteDisconnectError"); + py::register_exception(m, "nixlCancelledError"); + py::register_exception(m, "nixlNoTelemetryError"); py::class_(m, "nixlXferDList") - .def(py::init(), - py::arg("type"), - py::arg("sorted") = false, - py::arg("init_size") = 0) - .def(py::init([](nixl_mem_t mem, py::array descs, bool sorted) { + .def(py::init(), py::arg("type"), py::arg("init_size") = 0) + .def(py::init([](nixl_mem_t mem, py::array descs) { static_assert(sizeof(nixlBasicDesc) == 3 * sizeof(uint64_t), "nixlBasicDesc size mismatch"); // Check array shape and dtype @@ -190,20 +231,17 @@ PYBIND11_MODULE(_bindings, m) { throw std::invalid_argument("descs must be a C-contiguous numpy array"); } size_t n = descs.shape(0); - nixl_xfer_dlist_t new_list(mem, sorted, n); + nixl_xfer_dlist_t new_list(mem, n); // We assume that the Nx3 array matches the nixlBasicDesc layout so we can simply // memcpy std::memcpy(&new_list[0], descs.data(), descs.size() * sizeof(uint64_t)); - new_list.verifySorted(); - return new_list; }), py::arg("type"), - py::arg("descs").noconvert(), - py::arg("sorted") = false) - .def(py::init([](nixl_mem_t mem, py::list descs, bool sorted) { - nixl_xfer_dlist_t new_list(mem, sorted, descs.size()); + py::arg("descs").noconvert()) + .def(py::init([](nixl_mem_t mem, py::list descs) { + nixl_xfer_dlist_t new_list(mem, descs.size()); for (size_t i = 0; i < descs.size(); i++) { if (!py::isinstance(descs[i])) { throw py::type_error( @@ -219,17 +257,13 @@ PYBIND11_MODULE(_bindings, m) { desc[2].cast()); } - new_list.verifySorted(); - return new_list; }), py::arg("type"), - py::arg("descs").noconvert(), - py::arg("sorted") = false) + py::arg("descs").noconvert()) .def("getType", &nixl_xfer_dlist_t::getType) .def("descCount", &nixl_xfer_dlist_t::descCount) .def("isEmpty", &nixl_xfer_dlist_t::isEmpty) - .def("isSorted", &nixl_xfer_dlist_t::isSorted) .def(py::self == py::self) .def("__getitem__", [](nixl_xfer_dlist_t &list, unsigned int i) -> py::tuple { @@ -259,7 +293,6 @@ PYBIND11_MODULE(_bindings, m) { return (int)ret; }) .def("remDesc", &nixl_xfer_dlist_t::remDesc) - .def("verifySorted", &nixl_xfer_dlist_t::verifySorted) .def("clear", &nixl_xfer_dlist_t::clear) .def("print", &nixl_xfer_dlist_t::print) .def(py::pickle( @@ -276,11 +309,8 @@ PYBIND11_MODULE(_bindings, m) { })); py::class_(m, "nixlRegDList") - .def(py::init(), - py::arg("type"), - py::arg("sorted") = false, - py::arg("init_size") = 0) - .def(py::init([](nixl_mem_t mem, py::array descs, bool sorted) { + .def(py::init(), py::arg("type"), py::arg("init_size") = 0) + .def(py::init([](nixl_mem_t mem, py::array descs) { if (descs.ndim() != 2 || descs.shape(1) != 3) throw std::invalid_argument("descs must be a Nx3 numpy array"); if (!py::dtype::of().equal(descs.dtype()) && @@ -290,7 +320,7 @@ PYBIND11_MODULE(_bindings, m) { throw std::invalid_argument("descs must be a C-contiguous numpy array"); } size_t n = descs.shape(0); - nixl_reg_dlist_t new_list(mem, sorted, n); + nixl_reg_dlist_t new_list(mem, n); if (py::dtype::of().equal(descs.dtype())) { auto buffer = descs.unchecked(); for (size_t i = 0; i < n; i++) { @@ -303,12 +333,10 @@ PYBIND11_MODULE(_bindings, m) { } } - new_list.verifySorted(); - return new_list; })) - .def(py::init([](nixl_mem_t mem, py::list descs, bool sorted) { - nixl_reg_dlist_t new_list(mem, sorted, descs.size()); + .def(py::init([](nixl_mem_t mem, py::list descs) { + nixl_reg_dlist_t new_list(mem, descs.size()); for (size_t i = 0; i < descs.size(); i++) { if (!py::isinstance(descs[i])) { throw py::type_error( @@ -324,17 +352,14 @@ PYBIND11_MODULE(_bindings, m) { desc[2].cast(), desc[3].cast()); } - new_list.verifySorted(); return new_list; }), py::arg("type"), - py::arg("descs"), - py::arg("sorted") = false) + py::arg("descs")) .def("getType", &nixl_reg_dlist_t::getType) .def("descCount", &nixl_reg_dlist_t::descCount) .def("isEmpty", &nixl_reg_dlist_t::isEmpty) - .def("isSorted", &nixl_reg_dlist_t::isSorted) .def(py::self == py::self) .def("__getitem__", [](nixl_reg_dlist_t &list, unsigned int i) -> py::tuple { @@ -373,7 +398,6 @@ PYBIND11_MODULE(_bindings, m) { }) .def("trim", &nixl_reg_dlist_t::trim) .def("remDesc", &nixl_reg_dlist_t::remDesc) - .def("verifySorted", &nixl_reg_dlist_t::verifySorted) .def("clear", &nixl_reg_dlist_t::clear) .def("print", &nixl_reg_dlist_t::print) .def(py::pickle( @@ -394,7 +418,11 @@ PYBIND11_MODULE(_bindings, m) { .def(py::init()) .def(py::init()) .def(py::init()) - .def(py::init()); + .def(py::init()) + .def(py::init()) + .def(py::init()) + .def(py::init()) + .def(py::init()); // note: pybind will automatically convert notif_map to python types: // so, a Dictionary of string: List @@ -652,6 +680,15 @@ PYBIND11_MODULE(_bindings, m) { throw_nixl_exception(ret); return ret; }) + .def( + "getXferTelemetry", + [](nixlAgent &agent, uintptr_t reqh) -> nixl_xfer_telem_t { + nixl_xfer_telem_t telemetry; + nixl_status_t ret = agent.getXferTelemetry((nixlXferReqH *)reqh, telemetry); + throw_nixl_exception(ret); + return telemetry; + }, + py::arg("reqh")) .def("queryXferBackend", [](nixlAgent &agent, uintptr_t reqh) -> uintptr_t { nixlBackendH *backend = nullptr; diff --git a/src/bindings/rust/Cargo.toml b/src/bindings/rust/Cargo.toml index 8ad11c79fc..595403594e 100644 --- a/src/bindings/rust/Cargo.toml +++ b/src/bindings/rust/Cargo.toml @@ -15,16 +15,16 @@ [package] name = "nixl-sys" -version = "0.5.0" -edition = "2021" description = "Low-level bindings to the nixl library" -license = "Apache-2.0" -homepage = "https://github.com/ai-dynamo/nixl" -repository = "https://github.com/ai-dynamo/nixl.git" -authors = ["NIXL Developers "] -readme = "README.md" links = "nixl" build = "build.rs" +version.workspace = true +edition.workspace = true +authors.workspace = true +homepage.workspace = true +license.workspace = true +repository.workspace = true +readme.workspace = true [features] stub-api = [] @@ -33,7 +33,6 @@ stub-api = [] thiserror = { version = "2" } tracing = { version = "0.1" } serde = { version = "1", features = ["derive"] } - libc = "0.2" [build-dependencies] @@ -41,3 +40,6 @@ bindgen = "0.71" cc = { version = "1.2.23", features = ["parallel"] } pkg-config = "0.3" os_info = "3.11" + +[dev-dependencies] +tempfile = "3.20.0" diff --git a/src/bindings/rust/build.rs b/src/bindings/rust/build.rs index 89e0e2eb95..8134910fab 100644 --- a/src/bindings/rust/build.rs +++ b/src/bindings/rust/build.rs @@ -15,6 +15,7 @@ use std::env; use std::path::PathBuf; +use os_info; fn get_lib_path(nixl_root_path: &str, arch: &str) -> String { let os_info = os_info::get(); @@ -64,20 +65,69 @@ fn get_arch() -> String { } } +fn get_nixl_libs() -> Option> { + // Try to get all libraries, but return None if any fails + match ( + pkg_config::probe_library("nixl"), + pkg_config::probe_library("nixl_build"), + pkg_config::probe_library("nixl_common"), + pkg_config::probe_library("stream"), + pkg_config::probe_library("serdes"), + pkg_config::probe_library("ucx_utils"), + pkg_config::probe_library("etcd-cpp-api"), + pkg_config::probe_library("ucx"), + ) { + (Ok(nixl), Ok(nixl_build), Ok(nixl_common), Ok(stream), Ok(serdes), Ok(ucx_utils), Ok(etcd), Ok(ucx)) => { + Some(vec![nixl, nixl_build, nixl_common, stream, serdes, ucx_utils, etcd, ucx]) + } + _ => None, + } +} + fn build_nixl(cc_builder: &mut cc::Build) { let nixl_root_path = env::var("NIXL_PREFIX").unwrap_or_else(|_| "/opt/nvidia/nvda_nixl".to_string()); + + // Print the NIXL_PREFIX for debugging + println!("cargo:warning=Using NIXL_PREFIX: {}", nixl_root_path); + let nixl_include_path = format!("{}/include", nixl_root_path); + let nixl_include_paths = [ + &nixl_include_path, + "../../api/cpp", + "../../infra", + "../../core", + "/usr/include", + ]; + + let arch = get_arch(); + let nixl_lib_path = get_lib_path(&nixl_root_path, &arch); + + // Print the library path for debugging + println!("cargo:warning=Using library path: {}", nixl_lib_path); + + // Add all possible library paths + println!("cargo:rustc-link-search=native={}", nixl_lib_path); + println!("cargo:rustc-link-search=native={}/lib", nixl_root_path); + println!("cargo:rustc-link-search=native={}/lib64", nixl_root_path); + println!("cargo:rustc-link-search=native={}/lib/x86_64-linux-gnu", nixl_root_path); + + // Try to use pkg-config if available + if let Some(libs) = get_nixl_libs() { + println!("cargo:warning=Using pkg-config paths"); + for lib in libs { + for path in lib.link_paths { + println!("cargo:rustc-link-search=native={}", path.display()); + } + } + } else { + println!("cargo:warning=pkg-config not available, using manual library paths"); + } cc_builder .file("wrapper.cpp") - .include(&nixl_include_path) - .include("../../api/cpp") - .include("../../infra") - .include("../../core"); + .includes(nixl_include_paths); - let arch = get_arch(); - let nixl_lib_path = get_lib_path(&nixl_root_path, &arch); println!("cargo:rustc-link-search={}", nixl_lib_path); @@ -87,30 +137,54 @@ fn build_nixl(cc_builder: &mut cc::Build) { cc_builder.define("HAVE_ETCD", "1"); } - cc_builder.compile("wrapper"); + // Compile the wrapper C++ code + cc_builder.compile("nixl_wrapper"); - // Link against NIXL libraries in correct order - // Only link against etcd-cpp-api if it's enabled - if etcd_enabled { - println!("cargo:rustc-link-lib=dylib=etcd-cpp-api"); + // Get the output path for bindings + let out_path = PathBuf::from(env::var("OUT_DIR").unwrap()); + + // Generate bindings with minimal configuration + let mut builder = bindgen::Builder::default() + .header("wrapper.h") + .clang_arg("-std=c++17") + .clang_arg(format!("-I{}", nixl_include_path)) + .clang_arg("-I../../api/cpp") + .clang_arg("-I../../infra") + .clang_arg("-I../../core") + .clang_arg("-x") + .clang_arg("c++"); + + // Add system include paths if needed + if let Ok(cpp_include) = env::var("CPLUS_INCLUDE_PATH") { + for path in cpp_include.split(':') { + builder = builder.clang_arg(format!("-I{}", path)); + } } - println!("cargo:rustc-link-lib=dylib=stream"); - println!("cargo:rustc-link-lib=dylib=nixl_common"); + + // Link against required libraries + println!("cargo:rustc-link-lib=stdc++"); + + // Add NIXL libraries println!("cargo:rustc-link-lib=dylib=nixl"); println!("cargo:rustc-link-lib=dylib=nixl_build"); println!("cargo:rustc-link-lib=dylib=nixl_common"); - println!("cargo:rustc-link-lib=dylib=serdes"); - println!("cargo:rustc-link-lib=dylib=stream"); - println!("cargo:rustc-link-lib=dylib=ucx_utils"); - // Link against C++ standard library - println!("cargo:rustc-link-lib=dylib=stdc++"); + if etcd_enabled { + println!("cargo:rustc-link-lib=dylib=etcd-cpp-api"); + } // Tell cargo to invalidate the built crate whenever the wrapper changes println!("cargo:rustc-link-search=native={}", nixl_lib_path); println!("cargo:rerun-if-changed=wrapper.h"); println!("cargo:rerun-if-changed=wrapper.cpp"); println!("cargo:rerun-if-env-changed=HAVE_ETCD"); + + builder + .parse_callbacks(Box::new(bindgen::CargoCallbacks::new())) + .generate() + .expect("Unable to generate bindings") + .write_to_file(out_path.join("bindings.rs")) + .expect("Couldn't write bindings!"); } fn build_stubs(cc_builder: &mut cc::Build) { @@ -143,17 +217,6 @@ fn run_build(use_stub_api: bool) { } else { build_stubs(&mut cc_builder); } - - // Get the output path for bindings - let out_path = PathBuf::from(env::var("OUT_DIR").unwrap()); - - // Generate bindings - bindgen::Builder::default() - .header("wrapper.h") - .generate() - .expect("Unable to generate bindings") - .write_to_file(out_path.join("bindings.rs")) - .expect("Couldn't write bindings!"); } fn main() { diff --git a/src/bindings/rust/meson.build b/src/bindings/rust/meson.build index 759b6cfd6d..a60c6a4144 100644 --- a/src/bindings/rust/meson.build +++ b/src/bindings/rust/meson.build @@ -13,37 +13,44 @@ # See the License for the specific language governing permissions and # limitations under the License. -rustc = find_program('rustc', required: false) + cargo = find_program('cargo', required: false) +if cargo.found() and get_option('rust') + RUST_BINDINGS_NAME = 'nixl-sys' + + cargo_build_cmd = [ + cargo, 'build', '--target-dir', + 'src' / 'bindings' / 'rust' / RUST_BINDINGS_NAME, + ] + + if get_option('buildtype') == 'release' + cargo_build_cmd += ['--release'] + endif -nixl_include_dir = include_directories('../../api/cpp', '../../infra', '../../core') -nixl_dep = declare_dependency(link_with: nixl_lib, include_directories: nixl_include_dir) + nixl_include_dir = include_directories('../../api/cpp', '../../infra', '../../core') + nixl_dep = declare_dependency(link_with: nixl_lib, include_directories: nixl_include_dir) -wrapper_sources = ['wrapper.cpp'] -wrapper_lib = static_library('nixl_wrapper', - sources: wrapper_sources, - include_directories: nixl_include_dir, - link_with: nixl_lib, - dependencies: [nixl_dep] - ) + wrapper_sources = ['wrapper.cpp'] + wrapper_lib = static_library('nixl_wrapper', + sources: wrapper_sources, + include_directories: [nixl_include_dir, include_directories('/usr/include')], + link_with: nixl_lib, + dependencies: [nixl_dep] + ) -if rustc.found() and cargo.found() rust_lib = custom_target( - 'rust', - command: [ - cargo, 'build', '--release', - '--manifest-path', join_paths(meson.project_source_root(), 'src', 'bindings', 'rust', 'Cargo.toml') - ], - input: [], - output: 'libnixl-sys.so', - depends: [wrapper_lib], - install: false, - install_dir: get_option('libdir'), - env: [ - 'NIXL_PREFIX=' + get_option('prefix') - ] - ) + RUST_BINDINGS_NAME, + command: cargo_build_cmd, + input: [], + output: RUST_BINDINGS_NAME, + depends: [wrapper_lib], + install: true, + install_dir: '.', + env: [ + 'NIXL_PREFIX=' + get_option('prefix'), + ] + ) else nixl_rust_bindings_lib = disabler() endif diff --git a/src/bindings/rust/src/agent.rs b/src/bindings/rust/src/agent.rs index b06c6169f3..00ad05303b 100644 --- a/src/bindings/rust/src/agent.rs +++ b/src/bindings/rust/src/agent.rs @@ -14,6 +14,7 @@ // limitations under the License. use super::*; +use crate::descriptors::{QueryResponseList, RegDescList}; /// A NIXL agent that can create backends and manage memory #[derive(Debug, Clone)] @@ -21,6 +22,18 @@ pub struct Agent { inner: Arc>, } +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +pub enum XferStatus { + Success, + InProgress, +} + +impl XferStatus { + pub fn is_success(&self) -> bool { + return *self == XferStatus::Success; + } +} + impl Agent { /// Creates a new agent with the given name pub fn new(name: &str) -> Result { @@ -39,11 +52,11 @@ impl Agent { }) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(agent.name = %name, error = "invalid_param", "Failed to create NIXL agent"); + tracing::error!(agent.name = %name, error = "invalid_param", "Failed to create NIXL agent"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(agent.name = %name, error = "backend_error", "Failed to create NIXL agent"); + tracing::error!(agent.name = %name, error = "backend_error", "Failed to create NIXL agent"); Err(NixlError::BackendError) } } @@ -75,11 +88,11 @@ impl Agent { Ok(utils::StringList::new(inner)) } -1 => { - tracing::trace!(error = "invalid_param", "Failed to get NIXL plugins"); + tracing::error!(error = "invalid_param", "Failed to get NIXL plugins"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(error = "backend_error", "Failed to get NIXL plugins"); + tracing::error!(error = "backend_error", "Failed to get NIXL plugins"); Err(NixlError::BackendError) } } @@ -163,11 +176,11 @@ impl Agent { }) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(plugin.name = %plugin, error = "invalid_param", "Failed to create NIXL backend"); + tracing::error!(plugin.name = %plugin, error = "invalid_param", "Failed to create NIXL backend"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(plugin.name = %plugin, error = "backend_error", "Failed to create NIXL backend"); + tracing::error!(plugin.name = %plugin, error = "backend_error", "Failed to create NIXL backend"); Err(NixlError::BackendError) } } @@ -224,7 +237,7 @@ impl Agent { descriptor: &impl NixlDescriptor, opt_args: Option<&OptArgs>, ) -> Result { - let mut reg_dlist = RegDescList::new(descriptor.mem_type(), false)?; + let mut reg_dlist = RegDescList::new(descriptor.mem_type())?; unsafe { reg_dlist.add_storage_desc(descriptor)?; @@ -243,6 +256,41 @@ impl Agent { }) } + /// Query information about memory/storage + /// + /// # Arguments + /// * `descs` - Registration descriptor list to query + /// * `opt_args` - Optional arguments specifying backends + /// + /// # Returns + /// A list of query responses, where each response may contain parameters + /// describing the memory/storage characteristics. + pub fn query_mem( + &self, + descs: &RegDescList, + opt_args: Option<&OptArgs>, + ) -> Result { + let resp = QueryResponseList::new()?; + + let status = { + let inner_guard = self.inner.write().unwrap(); + unsafe { + nixl_capi_query_mem( + inner_guard.handle.as_ptr(), + descs.handle(), + resp.handle(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + } + }; + + match status { + NIXL_CAPI_SUCCESS => Ok(resp), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + /// Gets the local metadata for this agent as a byte array pub fn get_local_md(&self) -> Result, NixlError> { tracing::trace!("Getting local metadata"); @@ -279,11 +327,57 @@ impl Agent { Ok(bytes) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(error = "invalid_param", "Failed to get local metadata"); + tracing::error!(error = "invalid_param", "Failed to get local metadata"); + Err(NixlError::InvalidParam) + } + _ => { + tracing::error!(error = "backend_error", "Failed to get local metadata"); + Err(NixlError::BackendError) + } + } + } + + /// Gets the local partial metadata as a byte array + /// + /// # Arguments + /// * `descs` - Registration descriptor list to get metadata for + /// * `opt_args` - Optional arguments for getting metadata + /// + /// # Returns + /// A byte array containing the local partial metadata + /// + pub fn get_local_partial_md(&self, descs: &RegDescList, opt_args: Option<&OptArgs>) -> Result, NixlError> { + tracing::trace!("Getting local partial metadata"); + let mut data = std::ptr::null_mut(); + let mut len: usize = 0; + let inner_guard = self.inner.write().unwrap(); + + let status = unsafe { + nixl_capi_get_local_partial_md( + inner_guard.handle.as_ptr(), + descs.handle(), + &mut data as *mut *mut _, + &mut len, + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + match status { + NIXL_CAPI_SUCCESS => { + let bytes = unsafe { + let slice = std::slice::from_raw_parts(data as *const u8, len); + let vec = slice.to_vec(); + libc::free(data as *mut libc::c_void); + vec + }; + tracing::trace!(metadata.size = len, "Successfully retrieved local partial metadata"); + Ok(bytes) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!(error = "invalid_param", "Failed to get local partial metadata"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(error = "backend_error", "Failed to get local metadata"); + tracing::error!(error = "backend_error", "Failed to get local partial metadata"); Err(NixlError::BackendError) } } @@ -316,16 +410,141 @@ impl Agent { Ok(name) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(error = "invalid_param", "Failed to load remote metadata"); + tracing::error!(error = "invalid_param", "Failed to load remote metadata"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(error = "backend_error", "Failed to load remote metadata"); + tracing::error!(error = "backend_error", "Failed to load remote metadata"); Err(NixlError::BackendError) } } } + pub fn make_connection(&self, remote_agent: &str, opt_args: Option<&OptArgs>) -> Result<(), NixlError> { + let remote_agent = CString::new(remote_agent)?; + let inner_guard = self.inner.write().unwrap(); + + let status = unsafe { + nixl_capi_agent_make_connection( + inner_guard.handle.as_ptr(), + remote_agent.as_ptr(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + pub fn prepare_xfer_dlist( + &self, + agent_name: &str, + descs: &XferDescList, + opt_args: Option<&OptArgs>, + ) -> Result { + let c_agent_name = CString::new(agent_name)?; + let mut dlist_hndl = std::ptr::null_mut(); + let inner_guard = self.inner.read().unwrap(); + + let status = unsafe { + nixl_capi_prep_xfer_dlist( + inner_guard.handle.as_ptr(), + c_agent_name.as_ptr(), + descs.handle(), + &mut dlist_hndl, + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => Ok(XferDlistHandle::new(dlist_hndl, inner_guard.handle)), + _ => Err(NixlError::BackendError), + } + } + + pub fn make_xfer_req(&self, operation: XferOp, + local_descs: &XferDlistHandle, local_indices: &[i32], + remote_descs: &XferDlistHandle, remote_indices: &[i32], + opt_args: Option<&OptArgs>) -> Result { + let mut req = std::ptr::null_mut(); + let inner_guard = self.inner.read().unwrap(); + + let status = unsafe { + nixl_capi_make_xfer_req( + inner_guard.handle.as_ptr(), + operation as bindings::nixl_capi_xfer_op_t, + local_descs.handle(), + local_indices.as_ptr(), + local_indices.len() as usize, + remote_descs.handle(), + remote_indices.as_ptr(), + remote_indices.len() as usize, + &mut req, + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()) + ) + }; + + match status { + NIXL_CAPI_SUCCESS => Ok(XferRequest::new(NonNull::new(req) + .ok_or(NixlError::FailedToCreateXferRequest)?, + self.inner.clone(), + )), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Check if remote metadata for a specific agent is available + /// + /// This function checks if the metadata for the specified remote agent has been + /// loaded and if specific descriptors can be found in the metadata. + /// + /// # Arguments + /// * `remote_agent` - Name of the remote agent to check + /// * `descs` - Optional descriptor list to check against the remote metadata. + /// If None, only checks if any metadata exists for the agent. + /// + /// # Returns + /// `true` if the remote agent's metadata is available (and descriptors are found if provided), + /// `false` otherwise + pub fn check_remote_metadata(&self, remote_agent: &str, descs: Option<&XferDescList>) -> bool { + tracing::trace!(remote_agent = %remote_agent, "Checking remote metadata"); + + let c_remote_name = match CString::new(remote_agent) { + Ok(name) => name, + Err(_) => { + tracing::trace!( + error = "invalid_param", + remote_agent = %remote_agent, + "Invalid remote agent name" + ); + return false; + } + }; + + let status = unsafe { + bindings::nixl_capi_check_remote_md( + self.inner.read().unwrap().handle.as_ptr(), + c_remote_name.as_ptr(), + descs.map_or(std::ptr::null_mut(), |d| d.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => { + tracing::trace!(remote_agent = %remote_agent, "Remote metadata is available"); + true + } + _ => { + tracing::trace!(remote_agent = %remote_agent, "Remote metadata is not available"); + false + } + } + } + /// Invalidates a remote metadata for this agent pub fn invalidate_remote_md(&self, remote_agent: &str) -> Result<(), NixlError> { self.inner @@ -339,6 +558,216 @@ impl Agent { self.inner.write().unwrap().invalidate_all_remotes() } + /// Send this agent's metadata to etcdAdd commentMore actions + /// + /// This enables other agents to discover this agent's metadata via etcd. + /// + /// # Arguments + /// * `opt_args` - Optional arguments for sending metadata + pub fn send_local_md(&self, opt_args: Option<&OptArgs>) -> Result<(), NixlError> { + tracing::trace!("Sending local metadata to etcd"); + let inner_guard = self.inner.write().unwrap(); + let status = unsafe { + bindings::nixl_capi_send_local_md( + inner_guard.handle.as_ptr(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => { + tracing::trace!("Successfully sent local metadata to etcd"); + Ok(()) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!( + error = "invalid_param", + "Failed to send local metadata to etcd" + ); + Err(NixlError::InvalidParam) + } + _ => { + tracing::error!( + error = "backend_error", + "Failed to send local metadata to etcd" + ); + Err(NixlError::BackendError) + } + } + } + + /// Send this agent's partial metadata + /// + /// # Arguments + /// * `descs` - Registration descriptor list to send + /// * `opt_args` - Optional arguments for sending metadata + pub fn send_local_partial_md(&self, descs: &RegDescList, opt_args: Option<&OptArgs>) -> Result<(), NixlError> { + tracing::trace!("Sending local partial metadata to etcd"); + let inner_guard = self.inner.write().unwrap(); + let status = unsafe { + nixl_capi_send_local_partial_md( + inner_guard.handle.as_ptr(), + descs.handle(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + match status { + NIXL_CAPI_SUCCESS => { + tracing::trace!("Successfully sent local partial metadata to etcd"); + Ok(()) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!(error = "invalid_param", "Failed to send local partial metadata to etcd"); + Err(NixlError::InvalidParam) + } + _ => Err(NixlError::BackendError) + } + } + + + /// Fetch a remote agent's metadata from etcd + /// + /// Once fetched, the metadata will be loaded and cached locally, enabling + /// communication with the remote agent. + /// + /// # Arguments + /// * `remote_name` - Name of the remote agent to fetch metadata for + /// * `opt_args` - Optional arguments for fetching metadata + pub fn fetch_remote_md( + &self, + remote_name: &str, + opt_args: Option<&OptArgs>, + ) -> Result<(), NixlError> { + tracing::trace!(remote_agent = %remote_name, "Fetching remote metadata from etcd"); + + let c_remote_name = CString::new(remote_name)?; + let inner_guard = self.inner.write().unwrap(); + + let status = unsafe { + bindings::nixl_capi_fetch_remote_md( + inner_guard.handle.as_ptr(), + c_remote_name.as_ptr(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => { + self.inner + .write() + .unwrap() + .remotes + .insert(remote_name.to_string()); + tracing::trace!(remote_agent = %remote_name, "Successfully fetched remote metadata from etcd"); + Ok(()) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!(error = "invalid_param", remote_agent = %remote_name, "Failed to fetch remote metadata from etcd"); + Err(NixlError::InvalidParam) + } + _ => { + tracing::error!(error = "backend_error", remote_agent = %remote_name, "Failed to fetch remote metadata from etcd"); + Err(NixlError::BackendError) + } + } + } + + /// Invalidate this agent's metadata in etcd + /// + /// This signals to other agents that this agent's metadata is no longer valid. + /// + /// # Arguments + /// * `opt_args` - Optional arguments for invalidating metadata + pub fn invalidate_local_md(&self, opt_args: Option<&OptArgs>) -> Result<(), NixlError> { + tracing::trace!("Invalidating local metadata in etcd"); + let inner_guard = self.inner.write().unwrap(); + let status = unsafe { + bindings::nixl_capi_invalidate_local_md( + inner_guard.handle.as_ptr(), + opt_args.map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => { + tracing::trace!("Successfully invalidated local metadata in etcd"); + Ok(()) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!( + error = "invalid_param", + "Failed to invalidate local metadata in etcd" + ); + Err(NixlError::InvalidParam) + } + _ => { + tracing::error!( + error = "backend_error", + "Failed to invalidate local metadata in etcd" + ); + Err(NixlError::BackendError) + } + } + } + + /// Send a notification to a remote agent + /// + /// # Arguments + /// * `remote_agent` - Name of the remote agent to send notification to + /// * `message` - The notification message to send + /// * `backend` - Optional backend to use for sending the notification + /// + /// # Returns + /// `Ok(())` if the notification was sent successfully + pub fn send_notification( + &self, + remote_agent: &str, + message: &[u8], + backend: Option<&Backend>, + ) -> Result<(), NixlError> { + tracing::trace!(remote_agent = %remote_agent, "Sending notification"); + + let c_remote_name = CString::new(remote_agent)?; + let inner_guard = self.inner.write().unwrap(); + + let opt_args = if backend.is_some() { + let mut args = OptArgs::new()?; + if let Some(b) = backend { + args.add_backend(b)?; + } + Some(args) + } else { + None + }; + + let status = unsafe { + nixl_capi_gen_notif( + inner_guard.handle.as_ptr(), + c_remote_name.as_ptr(), + message.as_ptr() as *const std::ffi::c_void, + message.len(), + opt_args + .as_ref() + .map_or(std::ptr::null_mut(), |args| args.inner.as_ptr()), + ) + }; + + match status { + NIXL_CAPI_SUCCESS => { + tracing::trace!(remote_agent = %remote_agent, "Successfully sent notification"); + Ok(()) + } + NIXL_CAPI_ERROR_INVALID_PARAM => { + tracing::error!(error = "invalid_param", remote_agent = %remote_agent, "Failed to send notification"); + Err(NixlError::InvalidParam) + } + _ => { + tracing::error!(error = "backend_error", remote_agent = %remote_agent, "Failed to send notification"); + Err(NixlError::BackendError) + } + } + } + /// Creates a transfer request between local and remote descriptors /// /// # Arguments @@ -388,6 +817,44 @@ impl Agent { } } + /// Estimates the cost of a transfer request + /// + /// # Arguments + /// * `req` - Transfer request handle + /// * `opt_args` - Optional arguments for the estimation + /// + /// # Returns + /// A tuple containing (duration in microseconds, error margin in microseconds, cost method) + /// + /// # Errors + /// Returns a NixlError if the operation fails + pub fn estimate_xfer_cost( + &self, + req: &XferRequest, + opt_args: Option<&OptArgs>, + ) -> Result<(i64, i64, CostMethod), NixlError> { + let mut duration_us: i64 = 0; + let mut err_margin_us: i64 = 0; + let mut method: u32 = 0; + + let status = unsafe { + nixl_capi_estimate_xfer_cost( + self.inner.write().unwrap().handle.as_ptr(), + req.handle(), + opt_args.map_or(ptr::null_mut(), |args| args.inner.as_ptr()), + &mut duration_us, + &mut err_margin_us, + &mut method as *mut u32 as *mut bindings::nixl_capi_cost_t, + ) + }; + + match status { + NIXL_CAPI_SUCCESS => Ok((duration_us, err_margin_us, CostMethod::from(method))), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + /// Posts a transfer request to initiate a transfer /// /// After this, the transfer state can be checked asynchronously until completion. @@ -429,11 +896,11 @@ impl Agent { Ok(true) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(error = "invalid_param", "Failed to post transfer request"); + tracing::error!(error = "invalid_param", "Failed to post transfer request"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(error = "backend_error", "Failed to post transfer request"); + tracing::error!(error = "backend_error", "Failed to post transfer request"); Err(NixlError::BackendError) } } @@ -445,19 +912,49 @@ impl Agent { /// /// # Arguments /// * `req` - Transfer request handle after `post_xfer_req` - pub fn get_xfer_status(&self, req: &XferRequest) -> Result { + pub fn get_xfer_status(&self, req: &XferRequest) -> Result { let status = unsafe { nixl_capi_get_xfer_status(self.inner.write().unwrap().handle.as_ptr(), req.handle()) }; match status { - NIXL_CAPI_SUCCESS => Ok(false), // Transfer completed - NIXL_CAPI_IN_PROG => Ok(true), // Transfer in progress + NIXL_CAPI_SUCCESS => Ok(XferStatus::Success), // Transfer completed + NIXL_CAPI_IN_PROG => Ok(XferStatus::InProgress), // Transfer in progress NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), _ => Err(NixlError::BackendError), } } + /// Queries the backend for a transfer request + /// + /// # Arguments + /// * `req` - Transfer request handle after `post_xfer_req` + /// + /// # Returns + /// A handle to the backend used for the transfer + /// + /// # Errors + /// Returns a NixlError if the operation fails + pub fn query_xfer_backend(&self, req: &XferRequest) -> Result { + let mut backend = std::ptr::null_mut(); + let inner_guard = self.inner.write().unwrap(); + let status = unsafe { + nixl_capi_query_xfer_backend( + inner_guard.handle.as_ptr(), + req.handle(), + &mut backend + ) + }; + match status { + NIXL_CAPI_SUCCESS => { + Ok(Backend{ inner: NonNull::new(backend).ok_or(NixlError::FailedToCreateBackend)? }) + } + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Gets notifications from other agents /// /// # Arguments @@ -483,11 +980,11 @@ impl Agent { Ok(()) } NIXL_CAPI_ERROR_INVALID_PARAM => { - tracing::trace!(error = "invalid_param", "Failed to get notifications"); + tracing::error!(error = "invalid_param", "Failed to get notifications"); Err(NixlError::InvalidParam) } _ => { - tracing::trace!(error = "backend_error", "Failed to get notifications"); + tracing::error!(error = "backend_error", "Failed to get notifications"); Err(NixlError::BackendError) } } diff --git a/src/bindings/rust/src/descriptors.rs b/src/bindings/rust/src/descriptors.rs index 25e52940ab..a80206d95e 100644 --- a/src/bindings/rust/src/descriptors.rs +++ b/src/bindings/rust/src/descriptors.rs @@ -15,11 +15,15 @@ use super::*; +mod query; mod reg; mod xfer; +mod xfer_dlist_handle; +pub use query::{QueryResponse, QueryResponseIterator, QueryResponseList}; pub use reg::RegDescList; pub use xfer::XferDescList; +pub use xfer_dlist_handle::XferDlistHandle; /// Memory types supported by NIXL #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] diff --git a/src/bindings/rust/src/descriptors/query.rs b/src/bindings/rust/src/descriptors/query.rs new file mode 100644 index 0000000000..c4354d94f5 --- /dev/null +++ b/src/bindings/rust/src/descriptors/query.rs @@ -0,0 +1,172 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use super::*; +use crate::Params; + +/// A safe wrapper around a NIXL query response list +pub struct QueryResponseList { + inner: NonNull, +} + +/// Represents a single query response which may or may not contain parameters +pub struct QueryResponse<'a> { + list: &'a QueryResponseList, + index: usize, +} + +impl QueryResponseList { + /// Creates a new empty query response list + pub fn new() -> Result { + let mut list = ptr::null_mut(); + let status = unsafe { nixl_capi_create_query_resp_list(&mut list) }; + + match status { + NIXL_CAPI_SUCCESS => { + // SAFETY: If status is NIXL_CAPI_SUCCESS, list is non-null + let inner = unsafe { NonNull::new_unchecked(list) }; + Ok(Self { inner }) + } + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Returns the number of responses in the list + pub fn len(&self) -> Result { + let mut size = 0; + let status = unsafe { nixl_capi_query_resp_list_size(self.inner.as_ptr(), &mut size) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(size), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Returns true if the list is empty + pub fn is_empty(&self) -> Result { + Ok(self.len()? == 0) + } + + /// Gets a query response at the given index + pub fn get(&self, index: usize) -> Result, NixlError> { + let size = self.len()?; + if index >= size { + return Err(NixlError::InvalidParam); + } + + Ok(QueryResponse { list: self, index }) + } + + /// Returns an iterator + pub fn iter(&self) -> Result, NixlError> { + Ok(QueryResponseIterator { + list: self, + index: 0, + len: self.len()?, + }) + } + + pub(crate) fn handle(&self) -> *mut bindings::nixl_capi_query_resp_list_s { + self.inner.as_ptr() + } +} + +impl<'a> QueryResponse<'a> { + /// Returns true if this response contains parameters + pub fn has_value(&self) -> Result { + let mut has_value = false; + let status = unsafe { + nixl_capi_query_resp_list_has_value( + self.list.inner.as_ptr(), + self.index, + &mut has_value, + ) + }; + + match status { + NIXL_CAPI_SUCCESS => Ok(has_value), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Gets the parameters if this response has a value + pub fn get_params(&self) -> Result, NixlError> { + if !self.has_value()? { + return Ok(None); + } + + let mut params = ptr::null_mut(); + let status = unsafe { + nixl_capi_query_resp_list_get_params(self.list.inner.as_ptr(), self.index, &mut params) + }; + + match status { + NIXL_CAPI_SUCCESS => { + // SAFETY: If status is NIXL_CAPI_SUCCESS, params is non-null + let inner = unsafe { NonNull::new_unchecked(params) }; + Ok(Some(Params::new(inner))) + } + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } +} + +/// An iterator over query responses +pub struct QueryResponseIterator<'a> { + list: &'a QueryResponseList, + index: usize, + len: usize, +} + +impl<'a> Iterator for QueryResponseIterator<'a> { + type Item = QueryResponse<'a>; + + fn next(&mut self) -> Option { + if self.index >= self.len { + None + } else { + let response = QueryResponse { + list: self.list, + index: self.index, + }; + self.index += 1; + Some(response) + } + } + + fn size_hint(&self) -> (usize, Option) { + let remaining = self.len - self.index; + (remaining, Some(remaining)) + } +} + +impl<'a> ExactSizeIterator for QueryResponseIterator<'a> { + fn len(&self) -> usize { + self.len - self.index + } +} + +impl Drop for QueryResponseList { + fn drop(&mut self) { + // SAFETY: self.inner is guaranteed to be valid by NonNull + unsafe { + nixl_capi_destroy_query_resp_list(self.inner.as_ptr()); + } + } +} diff --git a/src/bindings/rust/src/descriptors/reg.rs b/src/bindings/rust/src/descriptors/reg.rs index 9f3e77db43..0fe3a987ee 100644 --- a/src/bindings/rust/src/descriptors/reg.rs +++ b/src/bindings/rust/src/descriptors/reg.rs @@ -23,10 +23,10 @@ pub struct RegDescList<'a> { impl<'a> RegDescList<'a> { /// Creates a new registration descriptor list for the given memory type - pub fn new(mem_type: MemType, sorted: bool) -> Result { + pub fn new(mem_type: MemType) -> Result { let mut dlist = ptr::null_mut(); let status = unsafe { - nixl_capi_create_reg_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist, sorted) + nixl_capi_create_reg_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist) }; match status { @@ -46,10 +46,38 @@ impl<'a> RegDescList<'a> { } } + pub fn get_type(&self) -> Result { + let mut mem_type = 0; + let status = unsafe { nixl_capi_reg_dlist_get_type(self.inner.as_ptr(), &mut mem_type) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(MemType::from(mem_type)), + _ => Err(NixlError::BackendError), + } + } + /// Adds a descriptor to the list pub fn add_desc(&mut self, addr: usize, len: usize, dev_id: u64) -> Result<(), NixlError> { + self.add_desc_with_meta(addr, len, dev_id, &[]) + } + + /// Add a descriptor with metadata + pub fn add_desc_with_meta( + &mut self, + addr: usize, + len: usize, + dev_id: u64, + metadata: &[u8], + ) -> Result<(), NixlError> { let status = unsafe { - nixl_capi_reg_dlist_add_desc(self.inner.as_ptr(), addr as uintptr_t, len, dev_id) + nixl_capi_reg_dlist_add_desc( + self.inner.as_ptr(), + addr as uintptr_t, + len, + dev_id, + metadata.as_ptr() as *const std::ffi::c_void, + metadata.len(), + ) }; match status { @@ -64,6 +92,17 @@ impl<'a> RegDescList<'a> { Ok(self.len()? == 0) } + /// Returns the number of descriptors in the list + pub fn desc_count(&self) -> Result { + let mut count = 0; + let status = unsafe { nixl_capi_reg_dlist_desc_count(self.inner.as_ptr(), &mut count) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(count), + _ => Err(NixlError::BackendError), + } + } + /// Returns the number of descriptors in the list pub fn len(&self) -> Result { let mut len = 0; @@ -76,14 +115,34 @@ impl<'a> RegDescList<'a> { } } - /// Returns true if any descriptors in the list overlap - pub fn has_overlaps(&self) -> Result { - let mut has_overlaps = false; - let status = - unsafe { nixl_capi_reg_dlist_has_overlaps(self.inner.as_ptr(), &mut has_overlaps) }; + /// Trims the list to the given size + pub fn trim(&mut self) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_reg_dlist_trim(self.inner.as_ptr()) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Removes the descriptor at the given index + pub fn rem_desc(&mut self, index: i32) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_reg_dlist_rem_desc(self.inner.as_ptr(), index) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Prints the list contents + pub fn print(&self) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_reg_dlist_print(self.inner.as_ptr()) }; match status { - NIXL_CAPI_SUCCESS => Ok(has_overlaps), + NIXL_CAPI_SUCCESS => Ok(()), NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), _ => Err(NixlError::BackendError), } @@ -129,8 +188,7 @@ impl<'a> RegDescList<'a> { _ => Err(NixlError::BackendError), }?; if len > 0 { - // TODO: Add API to get descriptor memory type - MemType::Unknown + self.get_type()? } else { desc_mem_type } diff --git a/src/bindings/rust/src/descriptors/xfer.rs b/src/bindings/rust/src/descriptors/xfer.rs index 16ab6e41e4..9f3816f532 100644 --- a/src/bindings/rust/src/descriptors/xfer.rs +++ b/src/bindings/rust/src/descriptors/xfer.rs @@ -23,10 +23,10 @@ pub struct XferDescList<'a> { impl<'a> XferDescList<'a> { /// Creates a new transfer descriptor list for the given memory type - pub fn new(mem_type: MemType, sorted: bool) -> Result { + pub fn new(mem_type: MemType) -> Result { let mut dlist = ptr::null_mut(); let status = unsafe { - nixl_capi_create_xfer_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist, sorted) + nixl_capi_create_xfer_dlist(mem_type as nixl_capi_mem_type_t, &mut dlist) }; match status { @@ -43,6 +43,21 @@ impl<'a> XferDescList<'a> { } } + pub fn as_ptr(&self) -> *mut bindings::nixl_capi_xfer_dlist_s { + self.inner.as_ptr() + } + + /// Returns the memory type of the transfer descriptor list + pub fn get_type(&self) -> Result { + let mut mem_type = 0; + let status = unsafe { nixl_capi_xfer_dlist_get_type(self.inner.as_ptr(), &mut mem_type) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(MemType::from(mem_type)), + _ => Err(NixlError::BackendError), + } + } + /// Adds a descriptor to the list pub fn add_desc(&mut self, addr: usize, len: usize, dev_id: u64) -> Result<(), NixlError> { let status = unsafe { @@ -61,6 +76,17 @@ impl<'a> XferDescList<'a> { Ok(self.len()? == 0) } + /// Returns the number of descriptors in the list + pub fn desc_count(&self) -> Result { + let mut count = 0; + let status = unsafe { nixl_capi_xfer_dlist_desc_count(self.inner.as_ptr(), &mut count) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(count), + _ => Err(NixlError::BackendError), + } + } + /// Returns the number of descriptors in the list pub fn len(&self) -> Result { let mut len = 0; @@ -73,14 +99,23 @@ impl<'a> XferDescList<'a> { } } - /// Returns true if any descriptors in the list overlap - pub fn has_overlaps(&self) -> Result { - let mut has_overlaps = false; - let status = - unsafe { nixl_capi_xfer_dlist_has_overlaps(self.inner.as_ptr(), &mut has_overlaps) }; + /// Trims the list to the given size + pub fn trim(&mut self) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_xfer_dlist_trim(self.inner.as_ptr()) }; match status { - NIXL_CAPI_SUCCESS => Ok(has_overlaps), + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Removes the descriptor at the given index + pub fn rem_desc(&mut self, index: i32) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_xfer_dlist_rem_desc(self.inner.as_ptr(), index) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(()), NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), _ => Err(NixlError::BackendError), } @@ -97,6 +132,17 @@ impl<'a> XferDescList<'a> { } } + /// Prints the list contents + pub fn print(&self) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_xfer_dlist_print(self.inner.as_ptr()) }; + + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + /// Resizes the list to the given size pub fn resize(&mut self, new_size: usize) -> Result<(), NixlError> { let status = unsafe { nixl_capi_xfer_dlist_resize(self.inner.as_ptr(), new_size) }; @@ -129,8 +175,7 @@ impl<'a> XferDescList<'a> { _ => Err(NixlError::BackendError), }?; if len > 0 { - // TODO: Add API to get descriptor memory type - MemType::Unknown + self.get_type().unwrap() } else { desc_mem_type } diff --git a/src/bindings/rust/src/descriptors/xfer_dlist_handle.rs b/src/bindings/rust/src/descriptors/xfer_dlist_handle.rs new file mode 100644 index 0000000000..73229ad11a --- /dev/null +++ b/src/bindings/rust/src/descriptors/xfer_dlist_handle.rs @@ -0,0 +1,41 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use super::*; + +pub struct XferDlistHandle { + inner: *mut bindings::nixl_capi_xfer_dlist_handle_s, + agent: NonNull +} + +impl XferDlistHandle { + pub fn new(inner: *mut bindings::nixl_capi_xfer_dlist_handle_s, + agent: NonNull) -> Self { + Self { inner, agent } + } + + pub fn handle(&self) -> *mut bindings::nixl_capi_xfer_dlist_handle_s { + self.inner + } +} + +impl Drop for XferDlistHandle { + fn drop(&mut self) { + unsafe { + nixl_capi_release_xfer_dlist_handle(self.agent.as_ptr(), + self.handle()); + } + } +} \ No newline at end of file diff --git a/src/bindings/rust/src/lib.rs b/src/bindings/rust/src/lib.rs index eecc110653..21c5560810 100644 --- a/src/bindings/rust/src/lib.rs +++ b/src/bindings/rust/src/lib.rs @@ -58,10 +58,21 @@ use bindings::{ nixl_capi_opt_args_set_notif_msg, nixl_capi_opt_args_set_skip_desc_merge, nixl_capi_params_create_iterator, nixl_capi_params_destroy_iterator, nixl_capi_params_is_empty, nixl_capi_params_iterator_next, nixl_capi_post_xfer_req, nixl_capi_reg_dlist_add_desc, - nixl_capi_reg_dlist_clear, nixl_capi_reg_dlist_has_overlaps, nixl_capi_reg_dlist_len, + nixl_capi_reg_dlist_clear, nixl_capi_reg_dlist_len, nixl_capi_reg_dlist_resize, nixl_capi_register_mem, nixl_capi_string_list_get, nixl_capi_string_list_size, nixl_capi_xfer_dlist_add_desc, nixl_capi_xfer_dlist_clear, - nixl_capi_xfer_dlist_has_overlaps, nixl_capi_xfer_dlist_len, nixl_capi_xfer_dlist_resize, + nixl_capi_xfer_dlist_len, nixl_capi_xfer_dlist_resize, + nixl_capi_agent_make_connection, nixl_capi_reg_dlist_get_type, nixl_capi_reg_dlist_desc_count, + nixl_capi_reg_dlist_trim, nixl_capi_reg_dlist_rem_desc, nixl_capi_reg_dlist_print, + nixl_capi_xfer_dlist_get_type, nixl_capi_xfer_dlist_desc_count, + nixl_capi_xfer_dlist_trim, nixl_capi_xfer_dlist_rem_desc, + nixl_capi_xfer_dlist_print, nixl_capi_gen_notif, nixl_capi_estimate_xfer_cost, + nixl_capi_query_mem, nixl_capi_create_query_resp_list, nixl_capi_destroy_query_resp_list, + nixl_capi_query_resp_list_size, nixl_capi_query_resp_list_has_value, + nixl_capi_query_resp_list_get_params, nixl_capi_prep_xfer_dlist, nixl_capi_release_xfer_dlist_handle, + nixl_capi_make_xfer_req, nixl_capi_get_local_partial_md, + nixl_capi_send_local_partial_md, nixl_capi_query_xfer_backend, nixl_capi_opt_args_set_ip_addr, + nixl_capi_opt_args_set_port }; // Re-export status codes @@ -81,6 +92,7 @@ mod xfer; pub use agent::*; pub use descriptors::*; pub use notify::*; +pub use utils::*; pub use xfer::*; /// Errors that can occur when using NIXL @@ -102,6 +114,10 @@ pub enum NixlError { RegDescListCreationFailed, #[error("Failed to add registration descriptor")] RegDescAddFailed, + #[error("Failed to create XferDlistHandle")] + FailedToCreateXferDlistHandle, + #[error("Failed to create backend")] + FailedToCreateBackend, } /// A safe wrapper around NIXL memory list @@ -143,7 +159,7 @@ impl RegistrationHandle { mem_type = ?self.mem_type, "Deregistering memory" ); - let mut reg_dlist = RegDescList::new(self.mem_type, false)?; + let mut reg_dlist = RegDescList::new(self.mem_type)?; unsafe { reg_dlist.add_desc(self.ptr, self.size, self.dev_id)?; let _opt_args = OptArgs::new().unwrap(); @@ -183,6 +199,7 @@ pub struct Backend { unsafe impl Send for Backend {} unsafe impl Sync for Backend {} + /// A safe wrapper around NIXL optional arguments pub struct OptArgs { inner: NonNull, @@ -305,6 +322,29 @@ impl OptArgs { _ => Err(NixlError::BackendError), } } + + /// Set the IP address + /// used in sendLocalMD, fetchRemoteMD, invalidateLocalMD, sendLocalPartialMD. + pub fn set_ip_addr(&mut self, ip_addr: &str) -> Result<(), NixlError> { + let c_str = CString::new(ip_addr).expect("Failed to convert string to CString"); + let status = unsafe { nixl_capi_opt_args_set_ip_addr(self.inner.as_ptr(), c_str.as_ptr()) }; + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } + + /// Set the port + /// used in sendLocalMD, fetchRemoteMD, invalidateLocalMD, sendLocalPartialMD. + pub fn set_port(&mut self, port: u16) -> Result<(), NixlError> { + let status = unsafe { nixl_capi_opt_args_set_port(self.inner.as_ptr(), port) }; + match status { + NIXL_CAPI_SUCCESS => Ok(()), + NIXL_CAPI_ERROR_INVALID_PARAM => Err(NixlError::InvalidParam), + _ => Err(NixlError::BackendError), + } + } } impl Drop for OptArgs { diff --git a/src/bindings/rust/src/tests.rs b/src/bindings/rust/src/tests.rs deleted file mode 100644 index 330c7a929d..0000000000 --- a/src/bindings/rust/src/tests.rs +++ /dev/null @@ -1,1046 +0,0 @@ -// SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -// SPDX-License-Identifier: Apache-2.0 -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -//! Raw FFI bindings to the NIXL library -//! -//! This crate provides low-level bindings to the NIXL C++ library. -//! It is not meant to be used directly, but rather through the higher-level -//! `nixl` crate. - -#[cfg(test)] -mod unit_tests { - use crate::*; - use std::env; - use std::time::Duration; - - // Helper function to create an agent with error handling - fn create_test_agent(name: &str) -> Result { - Agent::new(name) - } - - // Helper function to find a plugin by name - fn find_plugin(plugins: &StringList, name: &str) -> Result { - plugins - .iter() - .filter_map(Result::ok) - .find(|&plugin| plugin == name) - .map(ToString::to_string) - .or_else(|| plugins.get(0).ok().map(ToString::to_string)) - .ok_or(NixlError::InvalidParam) - } - - #[test] - fn test_agent_creation() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - drop(agent); - } - - #[test] - fn test_agent_invalid_name() { - let result = Agent::new("test\0agent"); - assert!(matches!(result, Err(NixlError::StringConversionError(_)))); - } - - #[test] - fn test_get_available_plugins() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let plugins = agent - .get_available_plugins() - .expect("Failed to get plugins"); - - // Print available plugins - for plugin in plugins.iter().flatten() { - println!("Found plugin: {}", plugin); - } - } - - #[test] - fn test_get_plugin_params() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let (_mems, _params) = agent - .get_plugin_params("UCX") - .expect("Failed to get plugin params"); - // MemList and Params will be automatically dropped here - } - - #[test] - fn test_backend_creation() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let (_mems, params) = agent - .get_plugin_params("UCX") - .expect("Failed to get plugin params"); - let backend = agent - .create_backend("UCX", ¶ms) - .expect("Failed to create backend"); - - let mut opt_args = OptArgs::new().expect("Failed to create opt args"); - opt_args - .add_backend(&backend) - .expect("Failed to add backend"); - } - - #[test] - fn test_params_iteration() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let (mems, params) = agent - .get_plugin_params("UCX") - .expect("Failed to get plugin params"); - - println!("Parameters:"); - if !params.is_empty().unwrap() { - for param in params.iter().unwrap() { - let param = param.unwrap(); - println!(" {} = {}", param.key, param.value); - } - } else { - println!(" (empty)"); - } - - println!("Memory types:"); - if !mems.is_empty().unwrap() { - for mem_type in mems.iter() { - println!(" {}", mem_type.unwrap()); - } - } else { - println!(" (empty)"); - } - } - - #[test] - fn test_get_backend_params() -> Result<(), NixlError> { - let agent = create_test_agent("test_agent")?; - let plugins = agent.get_available_plugins()?; - - // Ensure we have at least one plugin - assert!(!plugins.is_empty()?); - - // Try UCX plugin first since it doesn't require GPU - let plugin_name = find_plugin(&plugins, "UCX")?; - let (_mems, params) = agent.get_plugin_params(&plugin_name)?; - let backend = agent.create_backend(&plugin_name, ¶ms)?; - - // Get backend params after initialization - let (backend_mems, backend_params) = agent.get_backend_params(&backend)?; - - // Print parameters using iterator - let param_iter = backend_params.iter()?; - for param in param_iter.flatten() { - println!("Backend param: {} = {}", param.key, param.value); - } - - // Print memory types - for mem_type in backend_mems.iter().flatten() { - println!("Backend memory type: {:?}", mem_type); - } - - Ok(()) - } - - #[test] - fn test_xfer_dlist() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - - // Add some descriptors - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.add_desc(0x2000, 0x200, 1).unwrap(); - - // Check length - assert_eq!(dlist.desc_count().unwrap(), 2); - - // Check overlaps - assert!(!dlist.has_overlaps().unwrap()); - - // Add overlapping descriptor - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); - - // Clear list - dlist.clear().unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - - // Resize list - dlist.resize(5).unwrap(); - - // add descriptors with overlaps - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); - } - - #[test] - fn test_reg_dlist() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - - // Add some descriptors - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.add_desc(0x2000, 0x200, 1).unwrap(); - - // Check length - assert_eq!(dlist.desc_count().unwrap(), 2); - - // Check overlaps - assert!(!dlist.has_overlaps().unwrap()); - - // Add overlapping descriptor - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); - - // Clear list - dlist.clear().unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - - // Resize list - dlist.resize(5).unwrap(); - } - - #[test] - fn test_storage_descriptor_lifetime() { - // Create storage that outlives the descriptor list - let storage = SystemStorage::new(1024).unwrap(); - - { - // Create a descriptor list with shorter lifetime - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_storage_desc(&storage).unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 1); - // dlist is dropped here, but storage is still valid - } - - // MemoryRegion is still valid here - assert_eq!(::size(&storage), 1024); - } - - #[test] - fn test_multiple_storage_descriptors() { - let storage1 = SystemStorage::new(1024).unwrap(); - let storage2 = SystemStorage::new(2048).unwrap(); - - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - - // Add multiple descriptors - dlist.add_storage_desc(&storage1).unwrap(); - dlist.add_storage_desc(&storage2).unwrap(); - - assert_eq!(dlist.desc_count().unwrap(), 2); - } - - #[test] - fn test_memory_registration() { - let agent = Agent::new("test_agent").unwrap(); - let mut storage = SystemStorage::new(1024).unwrap(); - // Register memory - storage.register(&agent, None).unwrap(); - - // Verify we can still access the memory - storage.memset(0xAA); - assert!(storage.as_slice().iter().all(|&x| x == 0xAA)); - } - - #[test] - fn test_registration_handle_drop() { - let agent = Agent::new("test_agent").unwrap(); - let mut storage = SystemStorage::new(1024).unwrap(); - - // Register memory - storage.register(&agent, None).unwrap(); - - // Drop the storage, which should trigger deregistration - drop(storage); - - // Create new storage to verify we can register again - let mut new_storage = SystemStorage::new(1024).unwrap(); - new_storage.register(&agent, None).unwrap(); - } - - #[test] - fn test_multiple_registrations() { - let agent = Agent::new("test_agent").unwrap(); - let mut storage1 = SystemStorage::new(1024).unwrap(); - let mut storage2 = SystemStorage::new(2048).unwrap(); - - // Register both storages - storage1.register(&agent, None).unwrap(); - storage2.register(&agent, None).unwrap(); - - // Verify we can still access both memories - storage1.memset(0xAA); - storage2.memset(0xBB); - assert!(storage1.as_slice().iter().all(|&x| x == 0xAA)); - assert!(storage2.as_slice().iter().all(|&x| x == 0xBB)); - } - - #[test] - fn test_make_connection_success() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - // This should succeed if the agent is valid and the backend is set up - let result = agent.make_connection("remote_agent"); - // Accept either Ok or a backend error if no real remote exists - assert!( - result.is_ok() || matches!(result, Err(NixlError::BackendError)), - "Expected Ok or BackendError, got: {:?}", - result - ); - } - - #[test] - fn test_make_connection_invalid_param() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - // Null bytes in the name should trigger InvalidParam or StringConversionError - let result = agent.make_connection("remote\0agent"); - assert!( - matches!(result, Err(NixlError::StringConversionError(_))) || - matches!(result, Err(NixlError::InvalidParam)), - "Expected StringConversionError or InvalidParam, got: {:?}", - result - ); - } - - #[test] - fn test_prep_xfer_dlist_success() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let descs = XferDescList::new(MemType::Dram).unwrap(); - let opt_args = OptArgs::new().unwrap(); - let mut handle = XferDescListHandle::new().unwrap(); - let result = agent.prep_xfer_dlist("remote_agent", &descs, &mut handle, &opt_args); - // Accept Ok or BackendError if no real remote exists - assert!( - result.is_ok() || matches!(result, Err(NixlError::BackendError)), - "Expected Ok or BackendError, got: {:?}", - result - ); - } - - #[test] - fn test_prep_xfer_dlist_invalid_param() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let descs = XferDescList::new(MemType::Dram).unwrap(); - let opt_args = OptArgs::new().unwrap(); - let mut handle = XferDescListHandle::new().unwrap(); - // Null byte in agent name should trigger error - let result = agent.prep_xfer_dlist("remote\0agent", &descs, &mut handle, &opt_args); - assert!( - matches!(result, Err(NixlError::StringConversionError(_))) || - matches!(result, Err(NixlError::InvalidParam)), - "Expected StringConversionError or InvalidParam, got: {:?}", - result - ); - } - - #[test] - fn test_make_xfer_req_success() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let local_descs = XferDescList::new(MemType::Dram).unwrap(); - let remote_descs = XferDescList::new(MemType::Dram).unwrap(); - let opt_args = OptArgs::new().unwrap(); - let result = agent.make_xfer_req( - XferOp::Read, - &local_descs, - &remote_descs, - "remote_agent", - &opt_args, - ); - // Accept Ok or BackendError if no real remote exists - assert!( - result.is_ok() || matches!(result, Err(NixlError::BackendError)), - "Expected Ok or BackendError, got: {:?}", - result.err() - ); - } - - #[test] - fn test_make_xfer_req_invalid_param() { - let agent = Agent::new("test_agent").expect("Failed to create agent"); - let local_descs = XferDescList::new(MemType::Dram).unwrap(); - let remote_descs = XferDescList::new(MemType::Dram).unwrap(); - let opt_args = OptArgs::new().unwrap(); - // Null byte in remote_agent should trigger error - let result = agent.make_xfer_req( - XferOp::Read, - &local_descs, - &remote_descs, - "remote\0agent", - &opt_args, - ); - assert!( - matches!(result, Err(NixlError::StringConversionError(_))) || - matches!(result, Err(NixlError::InvalidParam)), - "Expected StringConversionError or InvalidParam, got: {}", - result.err().unwrap() - ); - } - - #[test] - fn test_get_local_md() { - let agent = Agent::new("test_agent").unwrap(); - - // Get available plugins and print their names - let plugins = agent.get_available_plugins().unwrap(); - for plugin in plugins.iter().flatten() { - println!("Found plugin: {}", plugin); - } - - // Get plugin parameters for both agents - let (_mem_list, params) = agent.get_plugin_params("UCX").unwrap(); - - // Create backends for both agents - let backend1 = agent.create_backend("UCX", ¶ms).unwrap(); - - let md = agent.get_local_md().unwrap(); - - // Measure the size - let initial_size = md.len(); - println!("Local metadata size: {}", initial_size); - - let mut opt_args = OptArgs::new().unwrap(); - opt_args.add_backend(&backend1).unwrap(); - - let mut storages = Vec::new(); - - for _i in 0..10 { - // Register some memory regions - let mut storage = SystemStorage::new(1024).unwrap(); - storage.register(&agent, Some(&opt_args)).unwrap(); - storages.push(storage); - } - - let md = agent.get_local_md().unwrap(); - - // Measure the size again - let final_size = md.len(); - println!("Local metadata size: {}", final_size); - - // Check if the size has increased - assert!(final_size > initial_size); - } - - #[test] - fn test_metadata_exchange() { - // Create two agents - let agent2 = Agent::new("agent2").unwrap(); - let agent1 = Agent::new("agent1").unwrap(); - - // Get plugin parameters for both agents - let (_mem_list, params) = agent1.get_plugin_params("UCX").unwrap(); - - // Create backends for both agents - let _backend1 = agent1.create_backend("UCX", ¶ms).unwrap(); - let _backend2 = agent2.create_backend("UCX", ¶ms).unwrap(); - - // Get metadata from agent1 - let md = agent1.get_local_md().unwrap(); - - // Load metadata into agent2 - let remote_name = agent2.load_remote_md(&md).unwrap(); - assert_eq!(remote_name, "agent1"); - } - - #[test] - fn test_basic_agent_lifecycle() -> Result<(), NixlError> { - // Create agents - let agent2 = create_test_agent("A2")?; - let agent1 = create_test_agent("A1")?; - - // Print available plugins - let plugins = agent1.get_available_plugins()?; - for plugin in plugins.iter().flatten() { - println!("Found plugin: {}", plugin); - } - - // Setup UCX backends - let (_mem_list1, _params) = agent1.get_plugin_params("UCX")?; - let (_mem_list2, params) = agent2.get_plugin_params("UCX")?; - - let _backend1 = agent1.create_backend("UCX", ¶ms)?; - let _backend2 = agent2.create_backend("UCX", ¶ms)?; - - // Setup memory regions - let mut storage1 = SystemStorage::new(256)?; - let mut storage2 = SystemStorage::new(256)?; - - // Initialize memory patterns - storage1.memset(0xbb); - storage2.memset(0x00); - - // Verify initial memory patterns - assert!(storage1.as_slice().iter().all(|&x| x == 0xbb)); - assert!(storage2.as_slice().iter().all(|&x| x == 0x00)); - - // Create registration descriptor lists - storage1.register(&agent1, None).unwrap(); - storage2.register(&agent2, None).unwrap(); - - // Exchange metadata - let metadata = agent2.get_local_md()?; - let remote_name = agent1.load_remote_md(&metadata)?; - assert_eq!(remote_name, "A2"); - - // Setup transfer descriptors - let mut local_xfer_dlist = XferDescList::new(MemType::Dram)?; - let mut remote_xfer_dlist = XferDescList::new(MemType::Dram)?; - local_xfer_dlist.add_storage_desc(&storage1)?; - remote_xfer_dlist.add_storage_desc(&storage2)?; - - // Setup transfer arguments - let mut xfer_args = OptArgs::new()?; - xfer_args.set_has_notification(true)?; - xfer_args.set_notification_message(b"notification")?; - - // Create and post transfer request - let xfer_req = agent1.create_xfer_req( - XferOp::Write, - &local_xfer_dlist, - &remote_xfer_dlist, - &remote_name, - Some(&xfer_args), - )?; - - // Handle transfer request - if let Ok(status) = agent1.post_xfer_req(&xfer_req, None) { - println!("Transfer request posted with status: {}", status); - - if status { - // Wait for transfer completion with timeout - let timeout = Duration::from_secs(5); - let start = std::time::Instant::now(); - - while start.elapsed() < timeout { - match agent1.get_xfer_status(&xfer_req) { - Ok(false) => { - println!("Transfer completed"); - break; - } - Ok(true) => std::thread::sleep(Duration::from_millis(100)), - Err(e) => { - println!("Error getting transfer status: {:?}", e); - break; - } - } - } - - // Wait for notifications with timeout - let mut notifs = NotificationMap::new()?; - let start = std::time::Instant::now(); - - while start.elapsed() < timeout { - match agent2.get_notifications(&mut notifs, None) { - Ok(_) if !notifs.is_empty()? => { - println!("Got notifications"); - break; - } - Ok(_) => std::thread::sleep(Duration::from_millis(100)), - Err(e) => { - println!("Error getting notifications: {:?}", e); - break; - } - } - } - - // Verify notification if received - if !notifs.is_empty()? { - let mut agents = notifs.agents(); - if let Some(Ok(agent_name)) = agents.next() { - let mut notifications = notifs.get_notifications(agent_name)?; - if let Some(Ok(notif)) = notifications.next() { - assert_eq!(notif, b"notification"); - - // Verify transfer if completed - if !agent1.get_xfer_status(&xfer_req)? { - assert!(storage2.as_slice().iter().all(|&x| x == 0xbb)); - } - } - } - } - } - } - - // Verify source memory remains unchanged - assert!(storage1.as_slice().iter().all(|&x| x == 0xbb)); - - Ok(()) - } - - #[test] - fn test_etcd_metadata_exchange() -> Result<(), NixlError> { - // Check if NIXL_ETCD_ENDPOINTS env var is set to skip test if not - if env::var("NIXL_ETCD_ENDPOINTS").is_err() { - println!("Skipping etcd test - NIXL_ETCD_ENDPOINTS not set"); - return Ok(()); - } - - // Create two agents for metadata exchange - let agent1 = Agent::new("EtcdAgent1")?; - let agent2 = Agent::new("EtcdAgent2")?; - - // Get UCX backend to add to optional arguments - let plugins = agent1.get_available_plugins()?; - let plugin_name = find_plugin(&plugins, "UCX")?; - let (_mems, params) = agent1.get_plugin_params(&plugin_name)?; - let backend = agent1.create_backend(&plugin_name, ¶ms)?; - - // Create OptArgs with backend - let mut opt_args = OptArgs::new()?; - opt_args.add_backend(&backend)?; - - // Send agent1's metadata to etcd - agent1.send_local_md(Some(&opt_args))?; - println!("Successfully sent agent1 metadata to etcd"); - - // Fetch agent1's metadata from etcd with agent2 - agent2.fetch_remote_md("EtcdAgent1", Some(&opt_args))?; - println!("Successfully fetched agent1 metadata from etcd"); - - // Invalidate agent1's metadata in etcd - agent1.invalidate_local_md(Some(&opt_args))?; - println!("Successfully invalidated agent1 metadata in etcd"); - - Ok(()) - } - - #[test] - fn test_send_notification() -> Result<(), NixlError> { - // Create two agents for notification exchange - let agent1 = Agent::new("NotifSender")?; - let agent2 = Agent::new("NotifReceiver")?; - - // Set up backends for both agents - let (_mem_list, params) = agent1.get_plugin_params("UCX")?; - let backend1 = agent1.create_backend("UCX", ¶ms)?; - let backend2 = agent2.create_backend("UCX", ¶ms)?; - - // Exchange metadata - let metadata = agent2.get_local_md()?; - agent1.load_remote_md(&metadata)?; - - // Create notification message - let message = b"Test notification message"; - - // Send notification with no backend specified - agent1.send_notification("NotifReceiver", message, None)?; - - // Send notification with specific backend - agent1.send_notification("NotifReceiver", message, Some(&backend1))?; - - // Create a notification map to receive notifications - let mut notifs = NotificationMap::new()?; - - // Receive notifications without backend - agent2.get_notifications(&mut notifs, None)?; - - // Receive notifications with specific backend - let mut opt_args = OptArgs::new()?; - opt_args.add_backend(&backend2)?; - agent2.get_notifications(&mut notifs, Some(&opt_args))?; - - // Verify notification map contents - if !notifs.is_empty()? { - let mut agents = notifs.agents(); - - // Should have notifications from NotifSender - if let Some(Ok(agent_name)) = agents.next() { - assert_eq!(agent_name, "NotifSender"); - - // Verify notification content - let notifications = notifs.get_notifications(agent_name)?; - let notif_count = notifs.get_notifications_size(agent_name)?; - - // May have 1 or 2 notifications depending on whether both were processed - assert!(notif_count > 0, "Should have at least one notification"); - - // Check content of notification - for notification in notifications { - assert_eq!(notification?, message); - } - } - } - - Ok(()) - } - - #[test] - fn test_check_remote_metadata() { - // Create two agents - let agent1 = Agent::new("agent1").expect("Failed to create agent1"); - let agent2 = Agent::new("agent2").expect("Failed to create agent2"); - - // Set up backends for both agents (required before metadata operations) - let (_mem_list, params) = agent1 - .get_plugin_params("UCX") - .expect("Failed to get plugin params"); - let _backend1 = agent1 - .create_backend("UCX", ¶ms) - .expect("Failed to create backend for agent1"); - let _backend2 = agent2 - .create_backend("UCX", ¶ms) - .expect("Failed to create backend for agent2"); - - // Initially, agent1 should not have metadata for agent2 - assert!(!agent1.check_remote_metadata("agent2", None)); - - // Get and share metadata - let metadata = agent2.get_local_md().expect("Failed to get local metadata"); - agent1 - .load_remote_md(&metadata) - .expect("Failed to load remote metadata"); - - // Now agent1 should have metadata for agent2 - assert!(agent1.check_remote_metadata("agent2", None)); - - // Test with a descriptor list - let mut storage = SystemStorage::new(1024).expect("Failed to create storage"); - let opt_args = OptArgs::new().expect("Failed to create opt args"); - storage - .register(&agent2, Some(&opt_args)) - .expect("Failed to register memory"); - - // Create descriptor list with memory that exists in agent2 - let mem_type = MemType::Dram; - let mut xfer_desc_list = - XferDescList::new(mem_type).expect("Failed to create xfer desc list"); - xfer_desc_list - .add_desc( - unsafe { storage.as_ptr() } as usize, - storage.size(), - storage.device_id(), - ) - .expect("Failed to add descriptor"); - - // Update metadata after registration - let metadata = agent2 - .get_local_md() - .expect("Failed to get updated local metadata"); - agent1 - .load_remote_md(&metadata) - .expect("Failed to reload remote metadata"); - - // Check with descriptor list - should return true for valid descriptors - assert!(agent1.check_remote_metadata("agent2", Some(&xfer_desc_list))); - - // Create a descriptor list with invalid memory address - let mut invalid_desc_list = - XferDescList::new(mem_type).expect("Failed to create invalid desc list"); - invalid_desc_list - .add_desc(0xdeadbeef, 1024, 0) - .expect("Failed to add invalid descriptor"); - - // Check with invalid descriptor list - should return false - assert!(!agent1.check_remote_metadata("agent2", Some(&invalid_desc_list))); - - // Check with non-existent agent name - assert!(!agent1.check_remote_metadata("non_existent_agent", None)); - - // Check with invalid agent name (contains null byte) - // The function should return false rather than panic - let invalid_name = "invalid\0agent"; - assert!(!agent1.check_remote_metadata(invalid_name, None)); - } - - #[test] - fn test_xfer_desc_list_new_and_new_sorted() { - let dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.is_empty().unwrap()); - let dlist_sorted = XferDescList::new_sorted(MemType::Dram).unwrap(); - assert!(dlist_sorted.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_new_sorted_sortedness() { - let dlist = XferDescList::new_sorted(MemType::Dram).unwrap(); - let sorted = dlist.is_sorted().unwrap(); - assert!(sorted); - let dlist_unsorted = XferDescList::new(MemType::Dram).unwrap(); - let unsorted = dlist_unsorted.is_sorted().unwrap(); - assert!(!unsorted); - } - - #[test] - fn test_xfer_desc_list_get_type() { - let dlist = XferDescList::new(MemType::Vram).unwrap(); - assert_eq!(dlist.get_type().unwrap(), MemType::Vram); - } - - #[test] - fn test_xfer_desc_list_get_type_after_add() { - let mut dlist = XferDescList::new(MemType::Block).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert_eq!(dlist.get_type().unwrap(), MemType::Block); - } - - #[test] - fn test_xfer_desc_list_verify_sorted_true() { - let mut dlist = XferDescList::new_sorted(MemType::Dram).unwrap(); - - // list size should be at least 1 to be considered sorted - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.verify_sorted().unwrap()); - } - - #[test] - fn test_xfer_desc_list_verify_sorted_false() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x2000, 0x100, 0).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); // out of order - assert!(!dlist.verify_sorted().unwrap()); - } - - #[test] - fn test_xfer_desc_list_desc_count_basic() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 1); - } - - #[test] - fn test_xfer_desc_list_desc_count_after_clear() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.clear().unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - } - - #[test] - fn test_xfer_desc_list_is_empty_true() { - let dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_is_empty_false() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(!dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_is_sorted_true() { - let mut dlist = XferDescList::new_sorted(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.is_sorted().unwrap()); - } - - #[test] - fn test_xfer_desc_list_is_sorted_false() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x2000, 0x100, 0).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(!dlist.is_sorted().unwrap()); - } - - #[test] - fn test_xfer_desc_list_trim_basic() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.trim().unwrap(); - assert!(dlist.desc_count().unwrap() <= 1); - } - - #[test] - fn test_xfer_desc_list_trim_empty() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.trim().is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_rem_desc_basic() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.rem_desc(0).is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_rem_desc_out_of_bounds() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.rem_desc(0).is_err()); - } - - #[test] - fn test_xfer_desc_list_clear_basic() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.clear().unwrap(); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_clear_empty() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.clear().is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_xfer_desc_list_print_basic() { - let dlist = XferDescList::new(MemType::Dram).unwrap(); - assert!(dlist.print().is_ok()); - } - - #[test] - fn test_xfer_desc_list_print_after_add() { - let mut dlist = XferDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.print().is_ok()); - } - - // ----------- RegDescList API TESTS ----------- - - #[test] - fn test_reg_desc_list_new_and_new_sorted() { - let dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.is_empty().unwrap()); - let dlist_sorted = RegDescList::new_sorted(MemType::Dram).unwrap(); - assert!(dlist_sorted.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_new_sorted_sortedness() { - let dlist = RegDescList::new_sorted(MemType::Dram).unwrap(); - assert!(dlist.is_sorted().unwrap()); - let dlist_unsorted = RegDescList::new(MemType::Dram).unwrap(); - assert!(!dlist_unsorted.is_sorted().unwrap()); - } - - #[test] - fn test_reg_desc_list_get_type() { - let dlist = RegDescList::new(MemType::Vram).unwrap(); - assert_eq!(dlist.get_type().unwrap(), MemType::Vram); - } - - #[test] - fn test_reg_desc_list_get_type_after_add() { - let mut dlist = RegDescList::new(MemType::Block).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert_eq!(dlist.get_type().unwrap(), MemType::Block); - } - - #[test] - fn test_reg_desc_list_verify_sorted_true() { - let mut dlist = RegDescList::new_sorted(MemType::Dram).unwrap(); - - // list size should be at least 1 to be considered sorted - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.verify_sorted().unwrap()); - } - - #[test] - fn test_reg_desc_list_verify_sorted_false() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x2000, 0x100, 0).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); // out of order - assert!(!dlist.verify_sorted().unwrap()); - } - - #[test] - fn test_reg_desc_list_desc_count_basic() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 1); - } - - #[test] - fn test_reg_desc_list_desc_count_after_clear() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.clear().unwrap(); - assert_eq!(dlist.desc_count().unwrap(), 0); - } - - #[test] - fn test_reg_desc_list_is_empty_true() { - let dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_is_empty_false() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(!dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_is_sorted_true() { - let mut dlist = RegDescList::new_sorted(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.is_sorted().unwrap()); - } - - #[test] - fn test_reg_desc_list_is_sorted_false() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x2000, 0x100, 0).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(!dlist.is_sorted().unwrap()); - } - - #[test] - fn test_reg_desc_list_trim_basic() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.trim().unwrap(); - assert!(dlist.desc_count().unwrap() <= 1); - } - - #[test] - fn test_reg_desc_list_trim_empty() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.trim().is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_rem_desc_basic() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.rem_desc(0).is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_rem_desc_out_of_bounds() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.rem_desc(0).is_err()); - } - - #[test] - fn test_reg_desc_list_clear_basic() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.clear().unwrap(); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_clear_empty() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.clear().is_ok()); - assert!(dlist.is_empty().unwrap()); - } - - #[test] - fn test_reg_desc_list_print_basic() { - let dlist = RegDescList::new(MemType::Dram).unwrap(); - assert!(dlist.print().is_ok()); - } - - #[test] - fn test_reg_desc_list_print_after_add() { - let mut dlist = RegDescList::new(MemType::Dram).unwrap(); - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - assert!(dlist.print().is_ok()); - } -} diff --git a/src/bindings/rust/stubs.cpp b/src/bindings/rust/stubs.cpp index ffde62c5c6..c172c1e544 100644 --- a/src/bindings/rust/stubs.cpp +++ b/src/bindings/rust/stubs.cpp @@ -24,6 +24,7 @@ extern "C" { +// clang-format off // Internal struct definitions to match our opaque types // These are now stubs as their internal details are no longer used. struct nixl_capi_agent_s { /* empty */ }; @@ -37,6 +38,9 @@ struct nixl_capi_xfer_dlist_s { /* empty */ }; struct nixl_capi_reg_dlist_s { /* empty */ }; struct nixl_capi_xfer_req_s { /* empty */ }; struct nixl_capi_notif_map_s { /* empty */ }; +struct nixl_capi_xfer_dlist_handle_s { /* empty */ }; + +// clang-format on nixl_capi_status_t nixl_capi_stub_abort() @@ -64,12 +68,50 @@ nixl_capi_get_local_md(nixl_capi_agent_t agent, void** data, size_t* len) return nixl_capi_stub_abort(); } +nixl_capi_status_t +nixl_capi_get_local_partial_md(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + void **data, + size_t *len, + nixl_capi_opt_args_t opt_args) { + return nixl_capi_stub_abort(); +} + nixl_capi_status_t nixl_capi_load_remote_md(nixl_capi_agent_t agent, const void* data, size_t len, char** agent_name) { return nixl_capi_stub_abort(); } +nixl_capi_status_t +nixl_capi_prep_xfer_dlist(nixl_capi_agent_t agent, + const char *agent_name, + nixl_capi_xfer_dlist_t descs, + nixl_capi_xfer_dlist_handle_t *dlist_handle, + nixl_capi_opt_args_t opt_args) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_release_xfer_dlist_handle(nixl_capi_agent_t agent, + nixl_capi_xfer_dlist_handle_t dlist_handle) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_make_xfer_req(nixl_capi_agent_t agent, + nixl_capi_xfer_op_t operation, + nixl_capi_xfer_dlist_handle_t local_descs, + const int *local_indices, + size_t local_indices_count, + nixl_capi_xfer_dlist_handle_t remote_descs, + const int *remote_indices, + size_t remote_indices_count, + nixl_capi_xfer_req_t *req_hndl, + nixl_capi_opt_args_t opt_args) { + return nixl_capi_stub_abort(); +} + nixl_capi_status_t nixl_capi_invalidate_remote_md(nixl_capi_agent_t agent, const char* remote_agent) { @@ -186,6 +228,16 @@ nixl_capi_opt_args_get_skip_desc_merge(nixl_capi_opt_args_t args, bool* skip_mer return nixl_capi_stub_abort(); } +nixl_capi_status_t +nixl_capi_opt_args_set_ip_addr(nixl_capi_opt_args_t args, const char *ip_addr) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_opt_args_set_port(nixl_capi_opt_args_t args, uint16_t port) { + return nixl_capi_stub_abort(); +} + nixl_capi_status_t nixl_capi_params_is_empty(nixl_capi_params_t params, bool* is_empty) { @@ -243,9 +295,8 @@ nixl_capi_get_backend_params( // Transfer descriptor list functions nixl_capi_status_t -nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t* dlist, bool sorted) -{ - return nixl_capi_stub_abort(); +nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t *dlist) { + return nixl_capi_stub_abort(); } nixl_capi_status_t @@ -266,12 +317,6 @@ nixl_capi_xfer_dlist_len(nixl_capi_xfer_dlist_t dlist, size_t* len) return nixl_capi_stub_abort(); } -nixl_capi_status_t -nixl_capi_xfer_dlist_has_overlaps(nixl_capi_xfer_dlist_t dlist, bool* has_overlaps) -{ - return nixl_capi_stub_abort(); -} - nixl_capi_status_t nixl_capi_xfer_dlist_clear(nixl_capi_xfer_dlist_t dlist) { @@ -286,9 +331,8 @@ nixl_capi_xfer_dlist_resize(nixl_capi_xfer_dlist_t dlist, size_t new_size) // Registration descriptor list functions nixl_capi_status_t -nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t* dlist, bool sorted) -{ - return nixl_capi_stub_abort(); +nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t *dlist) { + return nixl_capi_stub_abort(); } nixl_capi_status_t @@ -298,9 +342,13 @@ nixl_capi_destroy_reg_dlist(nixl_capi_reg_dlist_t dlist) } nixl_capi_status_t -nixl_capi_reg_dlist_add_desc(nixl_capi_reg_dlist_t dlist, uintptr_t addr, size_t len, uint64_t dev_id) -{ - return nixl_capi_stub_abort(); +nixl_capi_reg_dlist_add_desc(nixl_capi_reg_dlist_t dlist, + uintptr_t addr, + size_t len, + uint64_t dev_id, + const void *metadata, + size_t metadata_len) { + return nixl_capi_stub_abort(); } nixl_capi_status_t @@ -309,12 +357,6 @@ nixl_capi_reg_dlist_len(nixl_capi_reg_dlist_t dlist, size_t* len) return nixl_capi_stub_abort(); } -nixl_capi_status_t -nixl_capi_reg_dlist_has_overlaps(nixl_capi_reg_dlist_t dlist, bool* has_overlaps) -{ - return nixl_capi_stub_abort(); -} - nixl_capi_status_t nixl_capi_reg_dlist_clear(nixl_capi_reg_dlist_t dlist) { @@ -362,6 +404,13 @@ nixl_capi_get_xfer_status(nixl_capi_agent_t agent, nixl_capi_xfer_req_t req_hndl return nixl_capi_stub_abort(); } +nixl_capi_status_t +nixl_capi_query_xfer_backend(nixl_capi_agent_t agent, + nixl_capi_xfer_req_t req_hndl, + nixl_capi_backend_t *backend) { + return nixl_capi_stub_abort(); +} + nixl_capi_status_t nixl_capi_destroy_xfer_req(nixl_capi_xfer_req_t req) { @@ -423,4 +472,41 @@ nixl_capi_notif_map_clear(nixl_capi_notif_map_t map) return nixl_capi_stub_abort(); } +nixl_capi_status_t +nixl_capi_create_query_resp_list(nixl_capi_query_resp_list_t *list) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_destroy_query_resp_list(nixl_capi_query_resp_list_t list) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_query_resp_list_size(nixl_capi_query_resp_list_t list, size_t *size) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_query_resp_list_has_value(nixl_capi_query_resp_list_t list, + size_t index, + bool *has_value) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_query_resp_list_get_params(nixl_capi_query_resp_list_t list, + size_t index, + nixl_capi_params_t *params) { + return nixl_capi_stub_abort(); +} + +nixl_capi_status_t +nixl_capi_query_mem(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + nixl_capi_query_resp_list_t resp, + nixl_capi_opt_args_t opt_args) { + return nixl_capi_stub_abort(); +} + } // extern "C" diff --git a/src/bindings/rust/tests/tests.rs b/src/bindings/rust/tests/tests.rs index abd0fb0153..a6e60d3deb 100644 --- a/src/bindings/rust/tests/tests.rs +++ b/src/bindings/rust/tests/tests.rs @@ -20,6 +20,111 @@ //! `nixl` crate. use nixl_sys::*; +use std::env; + +// Helper function to create an agent with error handling +fn create_test_agent(name: &str) -> Result { + Agent::new(name) +} + +fn setup_agent_with_backend(agent: &Agent) -> Result { + let plugins = agent.get_available_plugins().expect("Failed to get available plugins"); + let plugin_name = find_plugin(&plugins, "UCX").expect("Failed to find plugin"); + let (_mems, params) = agent.get_plugin_params(&plugin_name).expect("Failed to get plugin params"); + agent.create_backend(&plugin_name, ¶ms).expect("Failed to create backend"); + + let mut opt_args = OptArgs::new().expect("Failed to create opt args"); + let _ = opt_args.add_backend(&agent.get_backend("UCX").unwrap()); + + Ok(opt_args) +} + +fn create_agent_with_backend(name: &str) -> Result<(Agent, OptArgs), NixlError> { + let agent = Agent::new(name).expect("Failed to create agent"); + let plugins = agent.get_available_plugins().expect("Failed to get available plugins"); + let plugin_name = find_plugin(&plugins, "UCX").expect("Failed to find plugin"); + let (_mems, params) = agent.get_plugin_params(&plugin_name).expect("Failed to get plugin params"); + agent.create_backend(&plugin_name, ¶ms).expect("Failed to create backend"); + + let mut opt_args = OptArgs::new().expect("Failed to create opt args"); + let _ = opt_args.add_backend(&agent.get_backend("UCX").unwrap()); + + Ok((agent, opt_args)) +} + + +fn create_storage_list(agent: &Agent, opt_args: &OptArgs, size: usize) -> Vec { + let mut storage_list = Vec::new(); + for _ in 0..size { + let mut storage = SystemStorage::new(1024).unwrap(); + storage.register(agent, Some(opt_args)).expect("Failed to register storage memory"); + storage.memset(0); + agent.register_memory(&storage, Some(opt_args)).expect("Failed to register storage memory"); + storage_list.push(storage); + } + storage_list +} + +fn create_dlist<'a>(storage_list: &'a mut Vec) -> Result, NixlError> { + let mut dlist = XferDescList::new(MemType::Dram).expect("Failed to create XferDescList"); + for storage in storage_list.iter_mut() { + dlist.add_storage_desc(storage).expect(&format!("Failed to add storage descriptor for storage")); + } + Ok(dlist) +} + +fn exchange_metadata(agent1: &Agent, agent2: &Agent) -> Result<(), NixlError> { + let metadata1 = agent1.get_local_md().expect("Failed to get local metadata"); + let metadata2 = agent2.get_local_md().expect("Failed to get local metadata"); + agent1.load_remote_md(&metadata2).expect("Failed to load remote metadata"); + agent2.load_remote_md(&metadata1).expect("Failed to load remote metadata"); + Ok(()) +} + +// Helper function to find a plugin by name +fn find_plugin(plugins: &StringList, name: &str) -> Result { + plugins + .iter() + .filter_map(Result::ok) + .find(|&plugin| plugin == name) + .map(ToString::to_string) + .or_else(|| plugins.get(0).ok().map(ToString::to_string)) + .ok_or(NixlError::InvalidParam) +} + +/// Helper function to create and initialize a POSIX backend with optional arguments +/// Returns (backend, opt_args) if POSIX is available, or None if not available +fn create_posix_backend(agent: &Agent) -> Option<(Backend, OptArgs)> { + // Get available plugins - check if POSIX is available + let plugins = agent + .get_available_plugins() + .expect("Failed to get plugins"); + + if !plugins + .iter() + .any(|p| p.as_ref().map(|s| *s == "POSIX").unwrap_or(false)) + { + println!("POSIX plugin not available, skipping test"); + return None; + } + + // Get plugin parameters and create POSIX backend + let (_mems, params) = agent + .get_plugin_params("POSIX") + .expect("Failed to get POSIX plugin params"); + + let backend = agent + .create_backend("POSIX", ¶ms) + .expect("Failed to create POSIX backend"); + + // Create optional arguments with the backend + let mut opt_args = OptArgs::new().expect("Failed to create opt args"); + opt_args + .add_backend(&backend) + .expect("Failed to add backend"); + + Some((backend, opt_args)) +} #[test] fn test_agent_creation() { @@ -98,35 +203,65 @@ fn test_params_iteration() { } } +// #[test] +// fn test_get_backend_params() { +// let agent = Agent::new("test_agent").unwrap(); +// let plugins = agent.get_available_plugins().unwrap(); +// assert!(!plugins.is_empty().unwrap_or(false)); + +// let plugin_name = plugins.get(0).unwrap(); +// let (_mems, params) = agent.get_plugin_params(plugin_name).unwrap(); +// let backend = agent.create_backend(plugin_name, ¶ms).unwrap(); + +// // Get backend params after initialization +// let (backend_mems, backend_params) = agent.get_backend_params(&backend).unwrap(); + +// // Verify we can access the parameters +// let param_iter = backend_params.iter().unwrap(); +// for param in param_iter { +// let param = param.unwrap(); +// println!("Backend param: {} = {}", param.key, param.value); +// } + +// // Verify we can access the memory types +// for mem_type in backend_mems.iter() { +// println!("Backend memory type: {:?}", mem_type.unwrap()); +// } +// } + #[test] -fn test_get_backend_params() { - let agent = Agent::new("test_agent").unwrap(); - let plugins = agent.get_available_plugins().unwrap(); - assert!(!plugins.is_empty().unwrap_or(false)); +fn test_get_backend_params() -> Result<(), NixlError> { + let agent = create_test_agent("test_agent")?; + let plugins = agent.get_available_plugins()?; + + // Ensure we have at least one plugin + assert!(!plugins.is_empty()?); - let plugin_name = plugins.get(0).unwrap(); - let (_mems, params) = agent.get_plugin_params(plugin_name).unwrap(); - let backend = agent.create_backend(plugin_name, ¶ms).unwrap(); + // Try UCX plugin first since it doesn't require GPU + let plugin_name = find_plugin(&plugins, "UCX")?; + let (_mems, params) = agent.get_plugin_params(&plugin_name)?; + let backend = agent.create_backend(&plugin_name, ¶ms)?; // Get backend params after initialization - let (backend_mems, backend_params) = agent.get_backend_params(&backend).unwrap(); + let (backend_mems, backend_params) = agent.get_backend_params(&backend)?; - // Verify we can access the parameters - let param_iter = backend_params.iter().unwrap(); - for param in param_iter { - let param = param.unwrap(); + // Print parameters using iterator + let param_iter = backend_params.iter()?; + for param in param_iter.flatten() { println!("Backend param: {} = {}", param.key, param.value); } - // Verify we can access the memory types - for mem_type in backend_mems.iter() { - println!("Backend memory type: {:?}", mem_type.unwrap()); + // Print memory types + for mem_type in backend_mems.iter().flatten() { + println!("Backend memory type: {:?}", mem_type); } + + Ok(()) } #[test] fn test_xfer_dlist() { - let mut dlist = XferDescList::new(MemType::Dram, false).unwrap(); + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); // Add some descriptors dlist.add_desc(0x1000, 0x100, 0).unwrap(); @@ -135,29 +270,17 @@ fn test_xfer_dlist() { // Check length assert_eq!(dlist.len().unwrap(), 2); - // Check overlaps - assert!(!dlist.has_overlaps().unwrap()); - - // Add overlapping descriptor - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); - // Clear list dlist.clear().unwrap(); assert_eq!(dlist.len().unwrap(), 0); // Resize list dlist.resize(5).unwrap(); - - // add descriptors with overlaps - dlist.add_desc(0x1000, 0x100, 0).unwrap(); - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); } #[test] fn test_reg_dlist() { - let mut dlist = RegDescList::new(MemType::Dram, false).unwrap(); + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); // Add some descriptors dlist.add_desc(0x1000, 0x100, 0).unwrap(); @@ -166,13 +289,6 @@ fn test_reg_dlist() { // Check length assert_eq!(dlist.len().unwrap(), 2); - // Check overlaps - assert!(!dlist.has_overlaps().unwrap()); - - // Add overlapping descriptor - dlist.add_desc(0x1050, 0x100, 0).unwrap(); - assert!(dlist.has_overlaps().unwrap()); - // Clear list dlist.clear().unwrap(); assert_eq!(dlist.len().unwrap(), 0); @@ -188,7 +304,7 @@ fn test_storage_descriptor_lifetime() { { // Create a descriptor list with shorter lifetime - let mut dlist = XferDescList::new(MemType::Dram, false).unwrap(); + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); dlist.add_storage_desc(&storage).unwrap(); assert_eq!(dlist.len().unwrap(), 1); // dlist is dropped here, but storage is still valid @@ -203,7 +319,7 @@ fn test_multiple_storage_descriptors() { let storage1 = SystemStorage::new(1024).unwrap(); let storage2 = SystemStorage::new(2048).unwrap(); - let mut dlist = XferDescList::new(MemType::Dram, false).unwrap(); + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); // Add multiple descriptors dlist.add_storage_desc(&storage1).unwrap(); @@ -258,6 +374,38 @@ fn test_multiple_registrations() { assert!(storage2.as_slice().iter().all(|&x| x == 0xBB)); } +#[test] +fn test_make_connection_success() { + let agent = Agent::new("test_agent").expect("Failed to create agent"); + let remote_agent = Agent::new("remote_agent").expect("Failed to create remote agent"); + let opt_args = setup_agent_with_backend(&agent).expect("Failed to setup agent"); + let _opt_args_remote = setup_agent_with_backend(&remote_agent).expect("Failed to setup agent"); + + exchange_metadata(&agent, &remote_agent).expect("Failed to exchange metadata"); + + // This should succeed if the agent is valid and the backend is set up + let result = agent.make_connection(&remote_agent.name(), Some(&opt_args)); + + assert!( + result.is_ok(), + "Expected Ok got: {:?}", + result + ); +} + +#[test] +fn test_make_connection_invalid_param() { + let agent = Agent::new("test_agent").expect("Failed to create agent"); + // Null bytes in the name should trigger InvalidParam or StringConversionError + let result = agent.make_connection("remote\0agent", None); + assert!( + matches!(result, Err(NixlError::StringConversionError(_))) || + matches!(result, Err(NixlError::InvalidParam)), + "Expected StringConversionError or InvalidParam, got: {:?}", + result + ); +} + #[test] fn test_get_local_md() { let agent = Agent::new("test_agent").unwrap(); @@ -369,10 +517,10 @@ fn test_basic_agent_lifecycle() { let remote_name = agent1.load_remote_md(&metadata).unwrap(); assert_eq!(remote_name, "A2"); - let mut local_xfer_dlist = XferDescList::new(MemType::Dram, false).unwrap(); + let mut local_xfer_dlist = XferDescList::new(MemType::Dram).unwrap(); local_xfer_dlist.add_storage_desc(&storage1).unwrap(); - let mut remote_xfer_dlist = XferDescList::new(MemType::Dram, false).unwrap(); + let mut remote_xfer_dlist = XferDescList::new(MemType::Dram).unwrap(); remote_xfer_dlist.add_storage_desc(&storage2).unwrap(); let mut xfer_args = OptArgs::new().unwrap(); @@ -396,7 +544,7 @@ fn test_basic_agent_lifecycle() { loop { let status = agent1.get_xfer_status(&xfer_req).unwrap(); - if !status { + if status.is_success() { println!("Xfer req completed"); break; } else { @@ -432,3 +580,774 @@ fn test_basic_agent_lifecycle() { assert!(storage1.as_slice().iter().all(|&x| x == 0xbb)); assert!(storage2.as_slice().iter().all(|&x| x == 0xbb)); } + +#[test] +fn test_etcd_metadata_exchange() -> Result<(), NixlError> { + // Check if NIXL_ETCD_ENDPOINTS env var is set to skip test if not + if env::var("NIXL_ETCD_ENDPOINTS").is_err() { + println!("Skipping etcd test - NIXL_ETCD_ENDPOINTS not set"); + return Ok(()); + } + + // Create two agents for metadata exchange + let agent1 = Agent::new("EtcdAgent1")?; + let agent2 = Agent::new("EtcdAgent2")?; + + // Get UCX backend to add to optional arguments + let plugins = agent1.get_available_plugins()?; + let plugin_name = find_plugin(&plugins, "UCX")?; + let (_mems, params) = agent1.get_plugin_params(&plugin_name)?; + let backend = agent1.create_backend(&plugin_name, ¶ms)?; + + // Create OptArgs with backend + let mut opt_args = OptArgs::new()?; + opt_args.add_backend(&backend)?; + + // Send agent1's metadata to etcd + agent1.send_local_md(Some(&opt_args))?; + println!("Successfully sent agent1 metadata to etcd"); + + // Fetch agent1's metadata from etcd with agent2 + agent2.fetch_remote_md("EtcdAgent1", Some(&opt_args))?; + println!("Successfully fetched agent1 metadata from etcd"); + + // Invalidate agent1's metadata in etcd + agent1.invalidate_local_md(Some(&opt_args))?; + println!("Successfully invalidated agent1 metadata in etcd"); + + Ok(()) +} + +#[test] +fn test_send_notification() -> Result<(), NixlError> { + // Create two agents for notification exchange + let agent1 = Agent::new("NotifSender")?; + let agent2 = Agent::new("NotifReceiver")?; + + // Set up backends for both agents + let (_mem_list, params) = agent1.get_plugin_params("UCX")?; + let backend1 = agent1.create_backend("UCX", ¶ms)?; + let backend2 = agent2.create_backend("UCX", ¶ms)?; + + // Exchange metadata + let metadata = agent2.get_local_md()?; + agent1.load_remote_md(&metadata)?; + + // Create notification message + let message = b"Test notification message"; + + // Send notification with no backend specified + agent1.send_notification("NotifReceiver", message, None)?; + + // Send notification with specific backend + agent1.send_notification("NotifReceiver", message, Some(&backend1))?; + + // Create a notification map to receive notifications + let mut notifs = NotificationMap::new()?; + + // Receive notifications without backend + agent2.get_notifications(&mut notifs, None)?; + + // Receive notifications with specific backend + let mut opt_args = OptArgs::new()?; + opt_args.add_backend(&backend2)?; + agent2.get_notifications(&mut notifs, Some(&opt_args))?; + + // Verify notification map contents + if !notifs.is_empty()? { + let mut agents = notifs.agents(); + + // Should have notifications from NotifSender + if let Some(Ok(agent_name)) = agents.next() { + assert_eq!(agent_name, "NotifSender"); + + // Verify notification content + let notifications = notifs.get_notifications(agent_name)?; + let notif_count = notifs.get_notifications_size(agent_name)?; + + // May have 1 or 2 notifications depending on whether both were processed + assert!(notif_count > 0, "Should have at least one notification"); + + // Check content of notification + for notification in notifications { + assert_eq!(notification?, message); + } + } + } + + Ok(()) +} + +#[test] +fn test_check_remote_metadata() { + // Create two agents + let agent1 = Agent::new("agent1").expect("Failed to create agent1"); + let agent2 = Agent::new("agent2").expect("Failed to create agent2"); + + // Set up backends for both agents (required before metadata operations) + let (_mem_list, params) = agent1 + .get_plugin_params("UCX") + .expect("Failed to get plugin params"); + let _backend1 = agent1 + .create_backend("UCX", ¶ms) + .expect("Failed to create backend for agent1"); + let _backend2 = agent2 + .create_backend("UCX", ¶ms) + .expect("Failed to create backend for agent2"); + + // Initially, agent1 should not have metadata for agent2 + assert!(!agent1.check_remote_metadata("agent2", None)); + + // Get and share metadata + let metadata = agent2.get_local_md().expect("Failed to get local metadata"); + agent1 + .load_remote_md(&metadata) + .expect("Failed to load remote metadata"); + + // Now agent1 should have metadata for agent2 + assert!(agent1.check_remote_metadata("agent2", None)); + + // Test with a descriptor list + let mut storage = SystemStorage::new(1024).expect("Failed to create storage"); + let opt_args = OptArgs::new().expect("Failed to create opt args"); + storage + .register(&agent2, Some(&opt_args)) + .expect("Failed to register memory"); + + // Create descriptor list with memory that exists in agent2 + let mem_type = MemType::Dram; + let mut xfer_desc_list = + XferDescList::new(mem_type).expect("Failed to create xfer desc list"); + xfer_desc_list + .add_desc( + unsafe { storage.as_ptr() } as usize, + storage.size(), + storage.device_id(), + ) + .expect("Failed to add descriptor"); + + // Update metadata after registration + let metadata = agent2 + .get_local_md() + .expect("Failed to get updated local metadata"); + agent1 + .load_remote_md(&metadata) + .expect("Failed to reload remote metadata"); + + // Check with descriptor list - should return true for valid descriptors + assert!(agent1.check_remote_metadata("agent2", Some(&xfer_desc_list))); + + // Create a descriptor list with invalid memory address + let mut invalid_desc_list = + XferDescList::new(mem_type).expect("Failed to create invalid desc list"); + invalid_desc_list + .add_desc(0xdeadbeef, 1024, 0) + .expect("Failed to add invalid descriptor"); + + // Check with invalid descriptor list - should return false + assert!(!agent1.check_remote_metadata("agent2", Some(&invalid_desc_list))); + + // Check with non-existent agent name + assert!(!agent1.check_remote_metadata("non_existent_agent", None)); + + // Check with invalid agent name (contains null byte) + // The function should return false rather than panic + let invalid_name = "invalid\0agent"; + assert!(!agent1.check_remote_metadata(invalid_name, None)); +} + +#[test] +fn test_xfer_desc_list_new() { + let dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_get_type() { + let dlist = XferDescList::new(MemType::Vram).unwrap(); + assert_eq!(dlist.get_type().unwrap(), MemType::Vram); +} + +#[test] +fn test_xfer_desc_list_get_type_after_add() { + let mut dlist = XferDescList::new(MemType::Block).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert_eq!(dlist.get_type().unwrap(), MemType::Block); +} + +#[test] +fn test_xfer_desc_list_desc_count_basic() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 0); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 1); +} + +#[test] +fn test_xfer_desc_list_desc_count_after_clear() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.clear().unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 0); +} + +#[test] +fn test_xfer_desc_list_is_empty_true() { + let dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_is_empty_false() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(!dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_trim_basic() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.trim().unwrap(); + assert!(dlist.desc_count().unwrap() <= 1); +} + +#[test] +fn test_xfer_desc_list_trim_empty() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.trim().is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_rem_desc_basic() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(dlist.rem_desc(0).is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_rem_desc_out_of_bounds() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.rem_desc(0).is_err()); +} + +#[test] +fn test_xfer_desc_list_clear_basic() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.clear().unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_clear_empty() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.clear().is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_xfer_desc_list_print_basic() { + let dlist = XferDescList::new(MemType::Dram).unwrap(); + assert!(dlist.print().is_ok()); +} + +#[test] +fn test_xfer_desc_list_print_after_add() { + let mut dlist = XferDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(dlist.print().is_ok()); +} + +// ----------- RegDescList API TESTS ----------- + +#[test] +fn test_reg_desc_list_new() { + let dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_get_type() { + let dlist = RegDescList::new(MemType::Vram).unwrap(); + assert_eq!(dlist.get_type().unwrap(), MemType::Vram); +} + +#[test] +fn test_reg_desc_list_get_type_after_add() { + let mut dlist = RegDescList::new(MemType::Block).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert_eq!(dlist.get_type().unwrap(), MemType::Block); +} + +#[test] +fn test_reg_desc_list_desc_count_basic() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 0); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 1); +} + +#[test] +fn test_reg_desc_list_desc_count_after_clear() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.clear().unwrap(); + assert_eq!(dlist.desc_count().unwrap(), 0); +} + +#[test] +fn test_reg_desc_list_is_empty_true() { + let dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_is_empty_false() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(!dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_trim_basic() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.trim().unwrap(); + assert!(dlist.desc_count().unwrap() <= 1); +} + +#[test] +fn test_reg_desc_list_trim_empty() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.trim().is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_rem_desc_basic() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(dlist.rem_desc(0).is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_rem_desc_out_of_bounds() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.rem_desc(0).is_err()); +} + +#[test] +fn test_reg_desc_list_clear_basic() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + dlist.clear().unwrap(); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_clear_empty() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.clear().is_ok()); + assert!(dlist.is_empty().unwrap()); +} + +#[test] +fn test_reg_desc_list_print_basic() { + let dlist = RegDescList::new(MemType::Dram).unwrap(); + assert!(dlist.print().is_ok()); +} + +#[test] +fn test_reg_desc_list_print_after_add() { + let mut dlist = RegDescList::new(MemType::Dram).unwrap(); + dlist.add_desc(0x1000, 0x100, 0).unwrap(); + assert!(dlist.print().is_ok()); +} + +#[test] +fn test_query_mem_with_files() { + use std::fs::File; + use std::io::Write; + + // Constants + const DESCRIPTOR_ADDR: usize = 0; + const DESCRIPTOR_SIZE: usize = 1024; + const DESCRIPTOR_DEV_ID: u64 = 0; + const NUM_FILES_TO_CREATE: usize = 2; + const EXPECTED_NUM_RESPONSES: usize = 3; + + // Create a unique temporary directory for this test + let temp_dir = tempfile::tempdir().expect("Failed to create temporary directory"); + let temp_dir_path = temp_dir.path(); + + // Define test files + let test_files = vec![ + ("test_query_mem_rust_1.txt", "Test content for file 1"), + ("test_query_mem_rust_2.txt", "Test content for file 2"), + ("non_existent_file_rust.txt", ""), // This file won't be created + ]; + + // Create temporary test files + let mut file_paths: Vec<_> = test_files + .iter() + .take(NUM_FILES_TO_CREATE) // Only create the first two files + .map(|(filename, content)| { + let file_path = temp_dir_path.join(filename); + let mut file = + File::create(&file_path).expect(&format!("Failed to create {}", filename)); + writeln!(file, "{}", content).expect(&format!("Failed to write to {}", filename)); + file_path + }) + .collect(); + + // Add the non-existent file path + file_paths.push(temp_dir_path.join(test_files[2].0)); + + // Create agent + let agent = Agent::new("test_agent").expect("Failed to create agent"); + + // Create POSIX backend + let (_backend, opt_args) = match create_posix_backend(&agent) { + Some(result) => result, + None => return, + }; + + // Create descriptor list with existing and non-existing files + let mut descs = + RegDescList::new(MemType::File).expect("Failed to create descriptor list"); + + // Add blob descriptors with filenames as metadata + for (i, file_path) in file_paths.iter().enumerate() { + descs + .add_desc_with_meta( + DESCRIPTOR_ADDR, + DESCRIPTOR_SIZE, + DESCRIPTOR_DEV_ID, + file_path.to_string_lossy().as_bytes(), + ) + .expect(&format!("Failed to add descriptor for file {}", i + 1)); + } + + // Query memory + let resp = agent + .query_mem(&descs, Some(&opt_args)) + .expect("Failed to query mem"); + + // Verify results + assert_eq!( + resp.len().unwrap(), + EXPECTED_NUM_RESPONSES, + "Expected 3 responses" + ); + + // Check responses - current order: existing file, existing file, non-existent file + let responses: Vec<_> = resp.iter().unwrap().collect(); + + assert!(responses[0].has_value().unwrap(), "First file should exist"); + assert!( + responses[1].has_value().unwrap(), + "Second file should exist" + ); + assert!( + !responses[2].has_value().unwrap(), + "Third file should not exist" + ); + + // Print parameters for existing files + for (i, response) in responses.iter().enumerate() { + if response.has_value().unwrap() { + if let Some(params) = response.get_params().unwrap() { + println!("Parameters for response {}:", i); + for param in params.iter().unwrap() { + let param = param.unwrap(); + println!(" {} = {}", param.key, param.value); + // POSIX backend returns mtime and mode parameters + if param.key == "mtime" || param.key == "mode" { + assert!( + !param.value.is_empty(), + "Parameter value should not be empty" + ); + } + } + } + } + } +} + +#[test] +fn test_query_mem_empty_list() { + // Constants + const EXPECTED_EMPTY_RESPONSES: usize = 0; + + // Create agent + let agent = Agent::new("test_agent").expect("Failed to create agent"); + + // Create POSIX backend + let (_backend, opt_args) = match create_posix_backend(&agent) { + Some(result) => result, + None => return, + }; + + // Create empty descriptor list + let descs = RegDescList::new(MemType::File).expect("Failed to create descriptor list"); + + // Query memory with empty list + let resp = agent + .query_mem(&descs, Some(&opt_args)) + .expect("Failed to query mem"); + + // Verify results + let num_responses = resp.len().expect("Failed to get response count"); + assert_eq!( + num_responses, EXPECTED_EMPTY_RESPONSES, + "Expected 0 responses for empty descriptor list" + ); +} + +// Tests for prep_xfer_dlist API +#[test] +fn test_prep_xfer_dlist_success() { + const DLIST_SIZE: usize = 10; + + // 1. Create agents and backends + let (local_agent, opt_args) = create_agent_with_backend("local_agent").expect("Failed to create agent"); + let (remote_agent, _opt_args_remote) = create_agent_with_backend("remote_agent").expect("Failed to create agent"); + + // 2. Create memory regions and register them + let mut storage_list = create_storage_list(&local_agent, &opt_args, DLIST_SIZE); + + { + // 3. Create transfer descriptor list + let dlist = create_dlist(&mut storage_list).expect("Failed to create XferDescList"); + + // 4. Exchange metadata + exchange_metadata(&local_agent, &remote_agent).expect("Failed to exchange metadata"); + + // 5. Prepare transfer descriptor list + let result = local_agent.prepare_xfer_dlist("", &dlist, None); + assert!(result.is_ok(), "prepare_xfer_dlist failed with error: {:?}", result.err()); + } +} + +#[test] +fn test_prep_xfer_dlist_invalid_agent() { + const DLIST_SIZE: usize = 10; + + let agent = Agent::new("test_agent").expect("Failed to create agent"); + let opt_args = setup_agent_with_backend(&agent).expect("Failed to setup agent"); + let mut storage_list = create_storage_list(&agent, &opt_args, DLIST_SIZE); + { + let dlist = create_dlist(&mut storage_list).expect("Failed to create XferDescList"); + + // Try with invalid agent name + let result = agent.prepare_xfer_dlist("invalid_agent", &dlist, None); + + assert!( + result.is_err_and(|e| matches!(e, NixlError::BackendError)), + "Expected InvalidParam for invalid agent name" + ); + } +} + +// Tests for make_xfer_req API +#[test] +fn test_make_xfer_req_success() { + const DLIST_SIZE: usize = 10; + + let (local_agent, opt_args) = create_agent_with_backend("local_agent").expect("Failed to create agent"); + let (remote_agent, opt_args_remote) = create_agent_with_backend("remote_agent").expect("Failed to create agent"); + + // 2. Create memory regions and register them + let mut storage_list = create_storage_list(&local_agent, &opt_args, DLIST_SIZE); + let mut remote_storage_list = create_storage_list(&remote_agent, &opt_args_remote, DLIST_SIZE); + + { + let dlist = create_dlist(&mut storage_list).expect("Failed to create XferDescList"); + let remote_dlist = create_dlist(&mut remote_storage_list).expect("Failed to create XferDescList"); + + // 4. Exchange metadata + exchange_metadata(&local_agent, &remote_agent).expect("Failed to exchange metadata"); + + // Prepare descriptor list handles + let local_handle: XferDlistHandle = local_agent.prepare_xfer_dlist("", &dlist, Some(&opt_args)) + .expect("Failed to prepare local descriptor list"); + + let remote_handle: XferDlistHandle = local_agent.prepare_xfer_dlist(remote_agent.name().as_str(), &remote_dlist, Some(&opt_args)) + .expect("Failed to prepare local descriptor list"); + + // Create transfer request using prepared handles with indices + let local_indices = (0..DLIST_SIZE).step_by(2).map(|i| i as i32).collect::>(); + let remote_indices = (1..DLIST_SIZE).step_by(2).map(|i| i as i32).collect::>(); + let result = local_agent.make_xfer_req( + XferOp::Write, + &local_handle, + &local_indices, + &remote_handle, + &remote_indices, + Some(&opt_args) + ); + + assert!( + result.is_ok(), + "make_xfer_req failed with error: {:?}", result.err() + ); + } +} + +#[test] +fn test_make_xfer_req_invalid_indices() { + const DLIST_SIZE: usize = 10; + let (agent1, opt_args) = create_agent_with_backend("agent1").expect("Failed to create agent"); + let (agent2, opt_args_remote) = create_agent_with_backend("agent2").expect("Failed to create agent"); + + let mut storage_list = create_storage_list(&agent1, &opt_args, DLIST_SIZE); + let mut remote_storage_list = create_storage_list(&agent2, &opt_args_remote, DLIST_SIZE); + + { + let local_dlist = create_dlist(&mut storage_list).expect("Failed to create descriptor list"); + let remote_dlist = create_dlist(&mut remote_storage_list).expect("Failed to create descriptor list"); + + exchange_metadata(&agent1, &agent2).expect("Failed to exchange metadata"); + + // Prepare descriptor list handles + let local_handle = agent1.prepare_xfer_dlist("", &local_dlist, Some(&opt_args)) + .expect("Failed to prepare local descriptor list"); + let remote_handle = agent1.prepare_xfer_dlist(agent2.name().as_str(), &remote_dlist, Some(&opt_args)) + .expect("Failed to prepare remote descriptor list"); + + // Test with out-of-bounds indices (should fail) + let invalid_indices = [999i32]; // Index 999 doesn't exist + let result = agent1.make_xfer_req( + XferOp::Write, + &local_handle, + &invalid_indices, // Out-of-bounds local index + &remote_handle, + &invalid_indices, // Out-of-bounds remote index + None + ); + assert!(result.is_err_and(|e| matches!(e, NixlError::BackendError)), "Expected InvalidParam for out-of-bounds indices"); + } +} + +// Tests for get_local_partial_md API +#[test] +fn test_get_local_partial_md_success() { + let (agent, opt_args) = create_agent_with_backend("test_agent") + .expect("Failed to setup agent with backend"); + let _storage_list = create_storage_list(&agent, &opt_args, 10); + // Create a registration descriptor list + let mut reg_descs = RegDescList::new(MemType::Dram) + .expect("Failed to create registration descriptor list"); + reg_descs.add_desc(0x1000, 0x100, 0) + .expect("Failed to add descriptor"); + // Get local partial metadata + let result = agent.get_local_partial_md(®_descs, Some(&opt_args)); + // Should succeed and return metadata + match result { + Ok(metadata) => { + assert!(!metadata.is_empty(), "Metadata should not be empty"); + println!("Partial metadata size: {}", metadata.len()); + } + Err(e) => { + // May fail if no partial metadata exists yet, which is acceptable + assert!( + matches!(e, NixlError::BackendError) || matches!(e, NixlError::InvalidParam), + "Expected BackendError or InvalidParam, got: {:?}", e + ); + } + } +} + +#[test] +fn test_get_local_partial_md_empty_descs() { + let (agent, _) = create_agent_with_backend("test_agent") + .expect("Failed to setup agent with backend"); + // Create empty registration descriptor list + let reg_descs = RegDescList::new(MemType::Dram) + .expect("Failed to create registration descriptor list"); + // Try with empty descriptor list should succeed and return all available backends + let result = agent.get_local_partial_md(®_descs, None); + assert!( + result.is_ok(), + "get_local_partial_md should succeed with empty descriptor list" + ); +} + +// Tests for send_local_partial_md API +#[test] +fn test_send_local_partial_md_success() { + let (agent, opt_args) = create_agent_with_backend("test_agent") + .expect("Failed to setup agent with backend"); + let (agent2, opt_args2) = create_agent_with_backend("test_agent2") + .expect("Failed to setup agent with backend"); + let _storage_list = create_storage_list(&agent, &opt_args, 10); + let _remote_storage_list = create_storage_list(&agent2, &opt_args2, 10); + + // Create a registration descriptor list + let mut reg_descs = RegDescList::new(MemType::Dram) + .expect("Failed to create registration descriptor list"); + reg_descs.add_storage_desc(&_storage_list[0]).expect("Failed to add storage descriptor"); + + // Send local partial metadata + let mut opt_args_temp = OptArgs::new().expect("Failed to create opt args"); + opt_args_temp.set_ip_addr("127.0.0.1").expect("Failed to set ip address"); + let result = agent.send_local_partial_md(®_descs, Some(&opt_args_temp)); + + assert!( + result.is_ok(), + "send_local_partial_md should succeed" + ); +} + +// Tests for query_xfer_backend API +#[test] +fn test_query_xfer_backend_success() { + let (agent1, opt_args) = create_agent_with_backend("agent1").expect("Failed to create agent"); + let (agent2, opt_args_remote) = create_agent_with_backend("agent2").expect("Failed to create agent"); + // Create descriptor lists + let mut storage_list = create_storage_list(&agent1, &opt_args, 1); + let mut remote_storage_list = create_storage_list(&agent2, &opt_args_remote, 1); + { + let local_dlist = create_dlist(&mut storage_list).expect("Failed to create descriptor list"); + let remote_dlist = create_dlist(&mut remote_storage_list).expect("Failed to create descriptor list"); + exchange_metadata(&agent1, &agent2).expect("Failed to exchange metadata"); + // Create transfer request + let xfer_req = agent1.create_xfer_req( + XferOp::Write, + &local_dlist, + &remote_dlist, + "agent2", + None + ).expect("Failed to create transfer request"); + // Query which backend will be used for this transfer + let result: Result = agent1.query_xfer_backend(&xfer_req); + assert!(result.is_ok(), "query_xfer_backend failed with error: {:?}", result.err()); + let backend = result.unwrap(); + println!("Transfer will use backend: {:?}", backend); + } +} +#[test] +fn test_query_xfer_backend_invalid_request() { + let (agent1, opt_args) = create_agent_with_backend("agent1").expect("Failed to create agent"); + let (agent2, opt_args_remote) = create_agent_with_backend("agent2").expect("Failed to create agent"); + // Create descriptor lists + let mut storage_list = create_storage_list(&agent1, &opt_args, 1); + let mut remote_storage_list = create_storage_list(&agent2, &opt_args_remote, 1); + { + let local_dlist = create_dlist(&mut storage_list).expect("Failed to create descriptor list"); + let remote_dlist = create_dlist(&mut remote_storage_list).expect("Failed to create descriptor list"); + // Create transfer request with non-existent remote agent (should fail or succeed) + let xfer_req_result = agent1.create_xfer_req( + XferOp::Write, + &local_dlist, + &remote_dlist, + "non_existent_agent", + None + ); + assert!(xfer_req_result.is_err(), "Transfer request creation should fail for non-existent agent"); + } +} diff --git a/src/bindings/rust/wrapper.cpp b/src/bindings/rust/wrapper.cpp index 615ac026c9..14b7ba4a2d 100644 --- a/src/bindings/rust/wrapper.cpp +++ b/src/bindings/rust/wrapper.cpp @@ -16,8 +16,8 @@ */ #include "wrapper.h" -#include -#include +#include "nixl.h" +#include "nixl_types.h" #include #include @@ -67,9 +67,8 @@ struct nixl_capi_xfer_dlist_s { nixl_xfer_dlist_t* dlist; }; -// Internal struct for descriptor list handle struct nixl_capi_xfer_dlist_handle_s { - nixlDlistH* dlist; + nixlDlistH *handle; }; struct nixl_capi_reg_dlist_s { @@ -85,6 +84,10 @@ struct nixl_capi_notif_map_s { nixl_notifs_t notif_map; }; +struct nixl_capi_query_resp_list_s { + std::vector responses; +}; + nixl_capi_status_t nixl_capi_create_agent(const char* name, nixl_capi_agent_t* agent) { @@ -156,6 +159,36 @@ nixl_capi_get_local_md(nixl_capi_agent_t agent, void** data, size_t* len) } } +nixl_capi_status_t +nixl_capi_get_local_partial_md(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + void **data, + size_t *len, + nixl_capi_opt_args_t opt_args) { + if (!agent || !descs || !data || !len) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + try { + nixl_blob_t blob; + nixl_opt_args_t *args = opt_args ? &opt_args->args : nullptr; + nixl_status_t ret = agent->inner->getLocalPartialMD(*descs->dlist, blob, args); + if (ret != NIXL_SUCCESS) { + return NIXL_CAPI_ERROR_BACKEND; + } + // Allocate memory for the blob data + *data = malloc(blob.size()); + if (!*data) { + return NIXL_CAPI_ERROR_BACKEND; + } + // Copy the data + memcpy(*data, blob.data(), blob.size()); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + nixl_capi_status_t nixl_capi_load_remote_md(nixl_capi_agent_t agent, const void* data, size_t len, char** agent_name) { @@ -222,6 +255,23 @@ nixl_capi_send_local_md(nixl_capi_agent_t agent, nixl_capi_opt_args_t opt_args) } } +nixl_capi_status_t +nixl_capi_send_local_partial_md(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + nixl_capi_opt_args_t opt_args) { + if (!agent || !descs) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + try { + nixl_opt_args_t *args = opt_args ? &opt_args->args : nullptr; + nixl_status_t ret = agent->inner->sendLocalPartialMD(*descs->dlist, args); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + nixl_capi_status_t nixl_capi_fetch_remote_md(nixl_capi_agent_t agent, const char* remote_name, nixl_capi_opt_args_t opt_args) { @@ -266,12 +316,12 @@ nixl_capi_check_remote_md(nixl_capi_agent_t agent, const char* remote_name, nixl try { // If descs is null, create an empty descriptor list of DRAM type if (!descs) { - nixl_xfer_dlist_t empty_list(DRAM_SEG, true); - nixl_status_t ret = agent->inner->checkRemoteMD(remote_name, empty_list); - return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + nixl_xfer_dlist_t empty_list(DRAM_SEG); + nixl_status_t ret = agent->inner->checkRemoteMD(remote_name, empty_list); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; } else { - nixl_status_t ret = agent->inner->checkRemoteMD(remote_name, *descs->dlist); - return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + nixl_status_t ret = agent->inner->checkRemoteMD(remote_name, *descs->dlist); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; } } catch (...) { @@ -615,6 +665,31 @@ nixl_capi_opt_args_get_skip_desc_merge(nixl_capi_opt_args_t args, bool* skip_mer } } +nixl_capi_status_t +nixl_capi_opt_args_set_ip_addr(nixl_capi_opt_args_t args, const char *ip_addr) { + if (!args || !ip_addr) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + args->args.ipAddr.assign(ip_addr); + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +nixl_capi_status_t +nixl_capi_opt_args_set_port(nixl_capi_opt_args_t args, uint16_t port) { + if (!args) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + args->args.port = port; + return NIXL_CAPI_SUCCESS; +} + nixl_capi_status_t nixl_capi_params_is_empty(nixl_capi_params_t params, bool* is_empty) { @@ -805,21 +880,20 @@ nixl_capi_get_backend_params( // Transfer descriptor list functions nixl_capi_status_t -nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t* dlist, bool sorted) -{ - if (!dlist) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } +nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t *dlist) { + if (!dlist) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } - try { - auto d = new nixl_capi_xfer_dlist_s; - d->dlist = new nixl_xfer_dlist_t(static_cast(mem_type), sorted); - *dlist = d; - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } + try { + auto d = new nixl_capi_xfer_dlist_s; + d->dlist = new nixl_xfer_dlist_t(static_cast(mem_type)); + *dlist = d; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } } nixl_capi_status_t @@ -911,22 +985,6 @@ nixl_capi_xfer_dlist_is_empty(nixl_capi_xfer_dlist_t dlist, bool* is_empty) } } -nixl_capi_status_t -nixl_capi_xfer_dlist_is_sorted(nixl_capi_xfer_dlist_t dlist, bool* is_sorted) -{ - if (!dlist || !is_sorted) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *is_sorted = dlist->dlist->isSorted(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - nixl_capi_status_t nixl_capi_xfer_dlist_trim(nixl_capi_xfer_dlist_t dlist) { @@ -958,38 +1016,6 @@ nixl_capi_status_t nixl_capi_xfer_dlist_rem_desc(nixl_capi_xfer_dlist_t dlist, i } } -nixl_capi_status_t -nixl_capi_xfer_dlist_has_overlaps(nixl_capi_xfer_dlist_t dlist, bool* has_overlaps) -{ - if (!dlist || !has_overlaps) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *has_overlaps = dlist->dlist->hasOverlaps(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - -nixl_capi_status_t -nixl_capi_xfer_dlist_verify_sorted(nixl_capi_xfer_dlist_t dlist, bool* is_sorted) -{ - if (!dlist || !is_sorted) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *is_sorted = dlist->dlist->verifySorted(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - nixl_capi_status_t nixl_capi_xfer_dlist_clear(nixl_capi_xfer_dlist_t dlist) { @@ -1037,53 +1063,22 @@ nixl_capi_xfer_dlist_resize(nixl_capi_xfer_dlist_t dlist, size_t new_size) } } -nixl_capi_status_t nixl_capi_create_xfer_dlist_handle(nixl_capi_xfer_dlist_handle_t* handle) -{ - if (!handle) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *handle = new nixl_capi_xfer_dlist_handle_s; - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - -nixl_capi_status_t nixl_capi_destroy_xfer_dlist_handle(nixl_capi_xfer_dlist_handle_t handle) -{ - if (!handle) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - delete handle; - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - // Registration descriptor list functions nixl_capi_status_t -nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t* dlist, bool sorted) -{ - if (!dlist) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } +nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t *dlist) { + if (!dlist) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } - try { - auto d = new nixl_capi_reg_dlist_s; - d->dlist = new nixl_reg_dlist_t(static_cast(mem_type), sorted); - *dlist = d; - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } + try { + auto d = new nixl_capi_reg_dlist_s; + d->dlist = new nixl_reg_dlist_t(static_cast(mem_type)); + *dlist = d; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } } nixl_capi_status_t @@ -1120,41 +1115,33 @@ nixl_capi_reg_dlist_get_type(nixl_capi_reg_dlist_t dlist, nixl_capi_mem_type_t* } nixl_capi_status_t -nixl_capi_reg_dlist_verify_sorted(nixl_capi_reg_dlist_t dlist, bool* is_sorted) -{ - if (!dlist || !is_sorted) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *is_sorted = dlist->dlist->verifySorted(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - -nixl_capi_status_t -nixl_capi_reg_dlist_add_desc(nixl_capi_reg_dlist_t dlist, uintptr_t addr, size_t len, uint64_t dev_id) -{ - if (!dlist) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } +nixl_capi_reg_dlist_add_desc(nixl_capi_reg_dlist_t dlist, + uintptr_t addr, + size_t len, + uint64_t dev_id, + const void *metadata, + size_t metadata_len) { + if (!dlist) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } - try { - nixlBlobDesc desc(addr, len, dev_id); // Empty metadata - dlist->dlist->addDesc(desc); + try { + nixl_blob_t meta_blob; + if (metadata && metadata_len > 0) { + meta_blob.assign((const char *)metadata, metadata_len); + } + nixlBlobDesc desc(addr, len, dev_id, meta_blob); + dlist->dlist->addDesc(desc); #ifdef NIXL_DEBUG - printf("** Adding descriptor\n"); - dlist->dlist->print(); - printf("** Added descriptor\n"); + printf("** Adding descriptor\n"); + dlist->dlist->print(); + printf("** Added descriptor\n"); #endif - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } } nixl_capi_status_t @@ -1195,37 +1182,6 @@ nixl_capi_reg_dlist_is_empty(nixl_capi_reg_dlist_t dlist, bool* is_empty) } } -nixl_capi_status_t nixl_capi_reg_dlist_is_sorted(nixl_capi_reg_dlist_t dlist, bool* is_sorted) -{ - if (!dlist || !is_sorted) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *is_sorted = dlist->dlist->isSorted(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - -nixl_capi_status_t -nixl_capi_reg_dlist_has_overlaps(nixl_capi_reg_dlist_t dlist, bool* has_overlaps) -{ - if (!dlist || !has_overlaps) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - *has_overlaps = dlist->dlist->hasOverlaps(); - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } -} - nixl_capi_status_t nixl_capi_reg_dlist_trim(nixl_capi_reg_dlist_t dlist) { if (!dlist) { @@ -1365,60 +1321,84 @@ nixl_capi_status_t nixl_capi_agent_make_connection( } } -nixl_capi_status_t nixl_capi_agent_prep_xfer_dlist( - nixl_capi_agent_t agent, const char* agent_name, nixl_capi_xfer_dlist_t descs, - nixl_capi_xfer_dlist_handle_t handle, nixl_capi_opt_args_t opt_args) -{ - auto backends = opt_args->args.backends; - - nixl_opt_args_t extra_params; - - if (!agent || !agent_name || !descs) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - for (nixlBackendH* backend : backends) { - extra_params.backends.push_back(backend); +nixl_capi_status_t +nixl_capi_prep_xfer_dlist(nixl_capi_agent_t agent, + const char *agent_name, + nixl_capi_xfer_dlist_t descs, + nixl_capi_xfer_dlist_handle_t *dlist_handle, + nixl_capi_opt_args_t opt_args) { + if (!agent || !agent_name || !descs) { + return NIXL_CAPI_ERROR_INVALID_PARAM; } - nixl_status_t ret = agent->inner->prepXferDlist(std::string(agent_name), *descs->dlist, - handle->dlist, &extra_params); - return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } + try { + *dlist_handle = new nixl_capi_xfer_dlist_handle_s; + nixl_status_t ret = agent->inner->prepXferDlist(std::string(agent_name), + *descs->dlist, + (*dlist_handle)->handle, + opt_args ? &opt_args->args : nullptr); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } } +nixl_capi_status_t +nixl_capi_release_xfer_dlist_handle(nixl_capi_agent_t agent, + nixl_capi_xfer_dlist_handle_t dlist_handle) { + if (!agent || !dlist_handle) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } -nixl_capi_status_t nixl_capi_agent_make_xfer_req( - nixl_capi_agent_t agent, nixl_capi_xfer_op_t operation, nixl_capi_xfer_dlist_t local_descs, - nixl_capi_xfer_dlist_t remote_descs, const char* remote_agent, nixl_capi_xfer_req_t* req_hndl, - nixl_capi_opt_args_t opt_args) -{ - if (!agent || !local_descs || !remote_descs || !remote_agent || !req_hndl) { - return NIXL_CAPI_ERROR_INVALID_PARAM; - } - - try { - auto req = new nixl_capi_xfer_req_s; - nixl_status_t ret = agent->inner->createXferReq( - static_cast(operation), *local_descs->dlist, *remote_descs->dlist, - std::string(remote_agent), req->req, opt_args ? &opt_args->args : nullptr); + try { + nixl_status_t ret = agent->inner->releasedDlistH(dlist_handle->handle); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} - if (ret != NIXL_SUCCESS) { - delete req; - return NIXL_CAPI_ERROR_BACKEND; +nixl_capi_status_t +nixl_capi_make_xfer_req(nixl_capi_agent_t agent, + nixl_capi_xfer_op_t operation, + nixl_capi_xfer_dlist_handle_t local_descs, + const int *local_indices, + size_t local_indices_count, + nixl_capi_xfer_dlist_handle_t remote_descs, + const int *remote_indices, + size_t remote_indices_count, + nixl_capi_xfer_req_t *req_hndl, + nixl_capi_opt_args_t opt_args) { + if (!agent || !local_descs || !remote_descs || !req_hndl) { + return NIXL_CAPI_ERROR_INVALID_PARAM; } - *req_hndl = req; - return NIXL_CAPI_SUCCESS; - } - catch (...) { - return NIXL_CAPI_ERROR_BACKEND; - } + try { + auto req = new nixl_capi_xfer_req_s; + nixl_status_t ret = agent->inner->makeXferReq( + static_cast(operation), + local_descs->handle, + std::vector(local_indices, local_indices + local_indices_count), + remote_descs->handle, + std::vector(remote_indices, remote_indices + remote_indices_count), + req->req, + opt_args ? &opt_args->args : nullptr); + + if (ret != NIXL_SUCCESS) { + delete req; + return NIXL_CAPI_ERROR_BACKEND; + } + + *req_hndl = req; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } } + nixl_capi_status_t nixl_capi_create_xfer_req( nixl_capi_agent_t agent, nixl_capi_xfer_op_t operation, nixl_capi_xfer_dlist_t local_descs, @@ -1505,6 +1485,28 @@ nixl_capi_get_xfer_status(nixl_capi_agent_t agent, nixl_capi_xfer_req_t req_hndl } } +nixl_capi_status_t +nixl_capi_query_xfer_backend(nixl_capi_agent_t agent, + nixl_capi_xfer_req_t req_hndl, + nixl_capi_backend_t *backend) { + if (!agent || !req_hndl || !backend) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + try { + auto backend_handle = new nixl_capi_backend_s; + nixl_status_t ret = agent->inner->queryXferBackend(req_hndl->req, backend_handle->backend); + if (ret != NIXL_SUCCESS) { + delete backend_handle; + return NIXL_CAPI_ERROR_BACKEND; + } + *backend = backend_handle; + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + nixl_capi_status_t nixl_capi_destroy_xfer_req(nixl_capi_xfer_req_t req) { @@ -1715,4 +1717,111 @@ nixl_capi_notif_map_clear(nixl_capi_notif_map_t map) } } +// Query response list functions +nixl_capi_status_t +nixl_capi_create_query_resp_list(nixl_capi_query_resp_list_t *list) { + if (!list) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + auto resp_list = new nixl_capi_query_resp_list_s; + *list = resp_list; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +nixl_capi_status_t +nixl_capi_destroy_query_resp_list(nixl_capi_query_resp_list_t list) { + if (!list) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + delete list; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +nixl_capi_status_t +nixl_capi_query_resp_list_size(nixl_capi_query_resp_list_t list, size_t *size) { + if (!list || !size) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + *size = list->responses.size(); + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +nixl_capi_status_t +nixl_capi_query_resp_list_has_value(nixl_capi_query_resp_list_t list, + size_t index, + bool *has_value) { + if (!list || !has_value || index >= list->responses.size()) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + *has_value = list->responses[index].has_value(); + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +nixl_capi_status_t +nixl_capi_query_resp_list_get_params(nixl_capi_query_resp_list_t list, + size_t index, + nixl_capi_params_t *params) { + if (!list || !params || index >= list->responses.size()) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + if (!list->responses[index].has_value()) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + auto param_list = new nixl_capi_params_s; + param_list->params = list->responses[index].value(); + *params = param_list; + return NIXL_CAPI_SUCCESS; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + +// Query memory function +nixl_capi_status_t +nixl_capi_query_mem(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + nixl_capi_query_resp_list_t resp, + nixl_capi_opt_args_t opt_args) { + if (!agent || !descs || !resp) { + return NIXL_CAPI_ERROR_INVALID_PARAM; + } + + try { + nixl_opt_args_t *args = opt_args ? &opt_args->args : nullptr; + nixl_status_t ret = agent->inner->queryMem(*descs->dlist, resp->responses, args); + return ret == NIXL_SUCCESS ? NIXL_CAPI_SUCCESS : NIXL_CAPI_ERROR_BACKEND; + } + catch (...) { + return NIXL_CAPI_ERROR_BACKEND; + } +} + } // extern "C" diff --git a/src/bindings/rust/wrapper.h b/src/bindings/rust/wrapper.h index 50c5b49145..8bac7169ac 100644 --- a/src/bindings/rust/wrapper.h +++ b/src/bindings/rust/wrapper.h @@ -16,12 +16,16 @@ */ #pragma once -#include -#include -#include #ifdef __cplusplus +#include +#include +#include extern "C" { +#else +#include +#include +#include #endif // Status codes for our C API @@ -56,6 +60,7 @@ struct nixl_capi_xfer_dlist_handle_s; struct nixl_capi_reg_dlist_s; struct nixl_capi_xfer_req_s; struct nixl_capi_notif_map_s; +struct nixl_capi_query_resp_list_s; // Opaque handle types for C++ objects typedef struct nixl_capi_agent_s* nixl_capi_agent_t; @@ -66,10 +71,11 @@ typedef struct nixl_capi_backend_s* nixl_capi_backend_t; typedef struct nixl_capi_opt_args_s* nixl_capi_opt_args_t; typedef struct nixl_capi_param_iter_s* nixl_capi_param_iter_t; typedef struct nixl_capi_xfer_dlist_s* nixl_capi_xfer_dlist_t; -typedef struct nixl_capi_xfer_dlist_handle_s* nixl_capi_xfer_dlist_handle_t; +typedef struct nixl_capi_xfer_dlist_handle_s *nixl_capi_xfer_dlist_handle_t; typedef struct nixl_capi_reg_dlist_s* nixl_capi_reg_dlist_t; typedef struct nixl_capi_xfer_req_s* nixl_capi_xfer_req_t; typedef struct nixl_capi_notif_map_s* nixl_capi_notif_map_t; +typedef struct nixl_capi_query_resp_list_s *nixl_capi_query_resp_list_t; // Transfer request functions typedef enum { @@ -85,6 +91,14 @@ nixl_capi_status_t nixl_capi_destroy_agent(nixl_capi_agent_t agent); // Get local metadata as a byte array nixl_capi_status_t nixl_capi_get_local_md(nixl_capi_agent_t agent, void** data, size_t* len); +// Get local partial metadata as a byte array +nixl_capi_status_t +nixl_capi_get_local_partial_md(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + void **data, + size_t *len, + nixl_capi_opt_args_t opt_args); + // Load remote metadata from a byte array nixl_capi_status_t nixl_capi_load_remote_md(nixl_capi_agent_t agent, const void* data, size_t len, char** agent_name); @@ -100,6 +114,12 @@ nixl_capi_status_t nixl_capi_check_remote_md(nixl_capi_agent_t agent, const char // Send local metadata to etcd nixl_capi_status_t nixl_capi_send_local_md(nixl_capi_agent_t agent, nixl_capi_opt_args_t opt_args); +// Send local partial metadata to etcd +nixl_capi_status_t +nixl_capi_send_local_partial_md(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + nixl_capi_opt_args_t opt_args); + // Fetch remote metadata from etcd nixl_capi_status_t nixl_capi_fetch_remote_md(nixl_capi_agent_t agent, const char* remote_name, nixl_capi_opt_args_t opt_args); @@ -136,6 +156,10 @@ nixl_capi_status_t nixl_capi_opt_args_set_has_notif(nixl_capi_opt_args_t args, b nixl_capi_status_t nixl_capi_opt_args_get_has_notif(nixl_capi_opt_args_t args, bool* has_notif); nixl_capi_status_t nixl_capi_opt_args_set_skip_desc_merge(nixl_capi_opt_args_t args, bool skip_merge); nixl_capi_status_t nixl_capi_opt_args_get_skip_desc_merge(nixl_capi_opt_args_t args, bool* skip_merge); +nixl_capi_status_t +nixl_capi_opt_args_set_ip_addr(nixl_capi_opt_args_t args, const char *ip_addr); +nixl_capi_status_t +nixl_capi_opt_args_set_port(nixl_capi_opt_args_t args, uint16_t port); // Parameter access functions nixl_capi_status_t nixl_capi_params_is_empty(nixl_capi_params_t params, bool* is_empty); @@ -160,15 +184,28 @@ nixl_capi_status_t nixl_capi_deregister_mem( nixl_capi_status_t nixl_capi_agent_make_connection( nixl_capi_agent_t agent, const char* remote_agent, nixl_capi_opt_args_t opt_args); -nixl_capi_status_t nixl_capi_agent_prep_xfer_dlist( - nixl_capi_agent_t agent, const char* agent_name, nixl_capi_xfer_dlist_t descs, - nixl_capi_xfer_dlist_handle_t handle, nixl_capi_opt_args_t opt_args); - -nixl_capi_status_t nixl_capi_agent_make_xfer_req( - nixl_capi_agent_t agent, nixl_capi_xfer_op_t operation, nixl_capi_xfer_dlist_t local_descs, - nixl_capi_xfer_dlist_t remote_descs, const char* remote_agent, nixl_capi_xfer_req_t* req_hndl, - nixl_capi_opt_args_t opt_args); - +nixl_capi_status_t +nixl_capi_prep_xfer_dlist(nixl_capi_agent_t agent, + const char *agent_name, + nixl_capi_xfer_dlist_t descs, + nixl_capi_xfer_dlist_handle_t *dlist_hndl, + nixl_capi_opt_args_t opt_args); + +nixl_capi_status_t +nixl_capi_release_xfer_dlist_handle(nixl_capi_agent_t agent, + nixl_capi_xfer_dlist_handle_t dlist_handle); + +nixl_capi_status_t +nixl_capi_make_xfer_req(nixl_capi_agent_t agent, + nixl_capi_xfer_op_t operation, + nixl_capi_xfer_dlist_handle_t local_descs, + const int *local_indices, + size_t local_indices_count, + nixl_capi_xfer_dlist_handle_t remote_descs, + const int *remote_indices, + size_t remote_indices_count, + nixl_capi_xfer_req_t *req_hndl, + nixl_capi_opt_args_t opt_args); // Notification functions nixl_capi_status_t nixl_capi_get_notifs( @@ -201,12 +238,18 @@ nixl_capi_status_t nixl_capi_post_xfer_req( nixl_capi_status_t nixl_capi_get_xfer_status(nixl_capi_agent_t agent, nixl_capi_xfer_req_t req_hndl); +nixl_capi_status_t +nixl_capi_query_xfer_backend(nixl_capi_agent_t agent, + nixl_capi_xfer_req_t req_hndl, + nixl_capi_backend_t *backend); + nixl_capi_status_t nixl_capi_release_xfer_req(nixl_capi_agent_t agent, nixl_capi_xfer_req_t req); nixl_capi_status_t nixl_capi_destroy_xfer_req(nixl_capi_xfer_req_t req); // Descriptor list functions -nixl_capi_status_t nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t* dlist, bool sorted); +nixl_capi_status_t +nixl_capi_create_xfer_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_xfer_dlist_t *dlist); nixl_capi_status_t nixl_capi_destroy_xfer_dlist(nixl_capi_xfer_dlist_t dlist); nixl_capi_status_t nixl_capi_xfer_dlist_get_type(nixl_capi_xfer_dlist_t dlist, nixl_capi_mem_type_t* mem_type); nixl_capi_status_t nixl_capi_xfer_dlist_add_desc( @@ -214,32 +257,28 @@ nixl_capi_status_t nixl_capi_xfer_dlist_add_desc( nixl_capi_status_t nixl_capi_xfer_dlist_desc_count(nixl_capi_xfer_dlist_t dlist, size_t* count); nixl_capi_status_t nixl_capi_xfer_dlist_len(nixl_capi_xfer_dlist_t dlist, size_t* len); nixl_capi_status_t nixl_capi_xfer_dlist_is_empty(nixl_capi_xfer_dlist_t dlist, bool* is_empty); -nixl_capi_status_t nixl_capi_xfer_dlist_is_sorted(nixl_capi_xfer_dlist_t dlist, bool* is_sorted); nixl_capi_status_t nixl_capi_xfer_dlist_trim(nixl_capi_xfer_dlist_t dlist); nixl_capi_status_t nixl_capi_xfer_dlist_rem_desc(nixl_capi_xfer_dlist_t dlist, int index); -nixl_capi_status_t nixl_capi_xfer_dlist_has_overlaps(nixl_capi_xfer_dlist_t dlist, bool* has_overlaps); -nixl_capi_status_t nixl_capi_xfer_dlist_verify_sorted(nixl_capi_xfer_dlist_t dlist, bool *is_sorted); nixl_capi_status_t nixl_capi_xfer_dlist_clear(nixl_capi_xfer_dlist_t dlist); nixl_capi_status_t nixl_capi_xfer_dlist_resize(nixl_capi_xfer_dlist_t dlist, size_t new_size); nixl_capi_status_t nixl_capi_xfer_dlist_print(nixl_capi_xfer_dlist_t dlist); -// Descriptor list handle functions -nixl_capi_status_t nixl_capi_create_xfer_dlist_handle(nixl_capi_xfer_dlist_handle_t* handle); -nixl_capi_status_t nixl_capi_destroy_xfer_dlist_handle(nixl_capi_xfer_dlist_handle_t handle); - -nixl_capi_status_t nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t* dlist, bool sorted); +nixl_capi_status_t +nixl_capi_create_reg_dlist(nixl_capi_mem_type_t mem_type, nixl_capi_reg_dlist_t *dlist); nixl_capi_status_t nixl_capi_destroy_reg_dlist(nixl_capi_reg_dlist_t dlist); nixl_capi_status_t nixl_capi_reg_dlist_get_type(nixl_capi_reg_dlist_t dlist, nixl_capi_mem_type_t* mem_type); -nixl_capi_status_t nixl_capi_reg_dlist_verify_sorted(nixl_capi_reg_dlist_t dlist, bool *is_sorted); -nixl_capi_status_t nixl_capi_reg_dlist_add_desc( - nixl_capi_reg_dlist_t dlist, uintptr_t addr, size_t len, uint64_t dev_id); +nixl_capi_status_t +nixl_capi_reg_dlist_add_desc(nixl_capi_reg_dlist_t dlist, + uintptr_t addr, + size_t len, + uint64_t dev_id, + const void *metadata, + size_t metadata_len); nixl_capi_status_t nixl_capi_reg_dlist_len(nixl_capi_reg_dlist_t dlist, size_t* len); nixl_capi_status_t nixl_capi_reg_dlist_desc_count(nixl_capi_reg_dlist_t dlist, size_t* count); nixl_capi_status_t nixl_capi_reg_dlist_is_empty(nixl_capi_reg_dlist_t dlist, bool* is_empty); -nixl_capi_status_t nixl_capi_reg_dlist_is_sorted(nixl_capi_reg_dlist_t dlist, bool* is_sorted); nixl_capi_status_t nixl_capi_reg_dlist_trim(nixl_capi_reg_dlist_t dlist); nixl_capi_status_t nixl_capi_reg_dlist_rem_desc(nixl_capi_reg_dlist_t dlist, int index); -nixl_capi_status_t nixl_capi_reg_dlist_has_overlaps(nixl_capi_reg_dlist_t dlist, bool* has_overlaps); nixl_capi_status_t nixl_capi_reg_dlist_clear(nixl_capi_reg_dlist_t dlist); nixl_capi_status_t nixl_capi_reg_dlist_resize(nixl_capi_reg_dlist_t dlist, size_t new_size); nixl_capi_status_t nixl_capi_reg_dlist_print(nixl_capi_reg_dlist_t dlist); @@ -250,6 +289,29 @@ nixl_capi_status_t nixl_capi_notif_map_get_notif( nixl_capi_notif_map_t map, const char* agent_name, size_t index, const void** data, size_t* len); nixl_capi_status_t nixl_capi_notif_map_clear(nixl_capi_notif_map_t map); +// Query response list functions +nixl_capi_status_t +nixl_capi_create_query_resp_list(nixl_capi_query_resp_list_t *list); +nixl_capi_status_t +nixl_capi_destroy_query_resp_list(nixl_capi_query_resp_list_t list); +nixl_capi_status_t +nixl_capi_query_resp_list_size(nixl_capi_query_resp_list_t list, size_t *size); +nixl_capi_status_t +nixl_capi_query_resp_list_has_value(nixl_capi_query_resp_list_t list, + size_t index, + bool *has_value); +nixl_capi_status_t +nixl_capi_query_resp_list_get_params(nixl_capi_query_resp_list_t list, + size_t index, + nixl_capi_params_t *params); + +// Query memory function +nixl_capi_status_t +nixl_capi_query_mem(nixl_capi_agent_t agent, + nixl_capi_reg_dlist_t descs, + nixl_capi_query_resp_list_t resp, + nixl_capi_opt_args_t opt_args); + #ifdef __cplusplus } #endif diff --git a/src/core/agent_data.h b/src/core/agent_data.h index 933da681c4..1b62fe3567 100644 --- a/src/core/agent_data.h +++ b/src/core/agent_data.h @@ -19,9 +19,11 @@ #include "common/str_tools.h" #include "mem_section.h" +#include "telemetry.h" #include "stream/metadata_stream.h" #include "sync.h" + #if HAVE_ETCD #include @@ -76,6 +78,9 @@ class nixlAgentData { std::unordered_map backendHandles; std::unordered_map connMD; + // Bookkeeping from GPU request handles to backend engines + std::unordered_map gpuReqToEngine; + // Local section, and Remote sections and their available common backends nixlLocalSection* memorySection; @@ -93,15 +98,28 @@ class nixlAgentData { std::mutex commLock; bool commThreadStop; bool useEtcd; - + std::unique_ptr telemetry_; void commWorker(nixlAgent* myAgent); void enqueueCommWork(nixl_comm_req_t request); void getCommWork(std::vector &req_list); + nixl_status_t + loadConnInfo(const std::string &remote_name, + const nixl_backend_t &backend, + const nixl_blob_t &conn_info); + nixl_status_t + loadRemoteSections(const std::string &remote_name, nixlSerDes &sd); + nixl_status_t + invalidateRemoteData(const std::string &remote_name); public: nixlAgentData(const std::string &name, const nixlAgentConfig &cfg); ~nixlAgentData(); + inline void + addErrorTelemetry(nixl_status_t err_status) { + if (telemetry_) telemetry_->updateErrorCount(err_status); + } + friend class nixlAgent; }; @@ -120,7 +138,6 @@ class nixlBackendH { bool supportsRemote () const { return engine->supportsRemote(); } bool supportsLocal () const { return engine->supportsLocal (); } bool supportsNotif () const { return engine->supportsNotif (); } - bool supportsProgTh () const { return engine->supportsProgTh(); } friend class nixlAgentData; friend class nixlAgent; diff --git a/src/core/meson.build b/src/core/meson.build index a78a31f9f5..221d053cba 100644 --- a/src/core/meson.build +++ b/src/core/meson.build @@ -21,8 +21,12 @@ if etcd_dep.found() nixl_lib_deps += [ etcd_dep ] endif +if 'LIBFABRIC' in static_plugins + nixl_lib_deps += [ libfabric_backend_interface, cuda_dep ] +endif + if 'UCX' in static_plugins - nixl_lib_deps += [ ucx_backend_interface, cuda_dep ] + nixl_lib_deps += [ ucx_backend_interface, asio_dep, cuda_dep ] endif if 'UCX_MO' in static_plugins @@ -54,11 +58,13 @@ if libtransfer_engine.found() and not disable_mooncake_backend and 'Mooncake' in endif nixl_lib = library('nixl', + 'signalhandler.cpp', 'nixl_agent.cpp', 'nixl_plugin_manager.cpp', 'nixl_listener.cpp', + 'telemetry.cpp', include_directories: [ nixl_inc_dirs, utils_inc_dirs ], - link_args: ['-lstdc++fs'], + link_args: ['-lstdc++fs', '-lbacktrace'], dependencies: nixl_lib_deps, install: true) diff --git a/src/core/nixl_agent.cpp b/src/core/nixl_agent.cpp index 72460ba36a..43a7b04580 100644 --- a/src/core/nixl_agent.cpp +++ b/src/core/nixl_agent.cpp @@ -18,6 +18,8 @@ #include #include #include +#include + #include "nixl.h" #include "serdes/serdes.h" #include "backend/backend_engine.h" @@ -25,7 +27,12 @@ #include "agent_data.h" #include "plugin_manager.h" #include "common/nixl_log.h" +#include "common/operators.h" +#include "telemetry.h" +#include "telemetry_event.h" +constexpr char TELEMETRY_ENABLED_VAR[] = "NIXL_TELEMETRY_ENABLE"; +constexpr char TELEMETRY_DIR_VAR[] = "NIXL_TELEMETRY_DIR"; static const std::vector> illegal_plugin_combinations = { {"GDS", "GDS_MT"}, }; @@ -47,7 +54,8 @@ std::string nixlEnumStrings::xferOpStr (const nixl_xfer_op_t &op) { } -std::string nixlEnumStrings::statusStr (const nixl_status_t &status) { +std::string +nixlEnumStrings::statusStr(const nixl_status_t &status) { switch (status) { case NIXL_IN_PROG: return "NIXL_IN_PROG"; case NIXL_SUCCESS: return "NIXL_SUCCESS"; @@ -61,28 +69,47 @@ std::string nixlEnumStrings::statusStr (const nixl_status_t &status) { case NIXL_ERR_UNKNOWN: return "NIXL_ERR_UNKNOWN"; case NIXL_ERR_NOT_SUPPORTED: return "NIXL_ERR_NOT_SUPPORTED"; case NIXL_ERR_REMOTE_DISCONNECT: return "NIXL_ERR_REMOTE_DISCONNECT"; + case NIXL_ERR_CANCELED: + return "NIXL_ERR_CANCELED"; + case NIXL_ERR_NO_TELEMETRY: + return "NIXL_ERR_NO_TELEMETRY"; default: return "BAD_STATUS"; } } -/*** nixlXferReqH telemetry update method, used mainly in the nixlAgent ***/ -void -nixlXferReqH::updateRequestStats(const std::string &dbg_msg_type) { - const auto xfer_time = std::chrono::duration_cast( - std::chrono::high_resolution_clock::now() - telemetry.startTime); - // If endTime needs to be recorded per Xfer, now() value here can be returned - - // To be replaced with NIXL_DEBUG when full telemetry is added - std::cout << "[NIXL TELEMETRY]: From backend " << engine->getType() << " " << dbg_msg_type - << " Xfer with " << initiatorDescs->descCount() << " descriptors of total size " - << telemetry.totalBytes << "B in " << xfer_time.count() << "us." << std::endl; +inline void +nixlXferReqH::updateRequestStats(std::unique_ptr &telemetry_pub, + nixl_telemetry_stat_status_t stat_status) { + + static const std::array nixl_post_status_str = { + " Posted", " Posted and Completed", " Completed"}; + auto duration = std::chrono::duration_cast( + std::chrono::steady_clock::now() - telemetry.startTime); + if (stat_status == NIXL_TELEMETRY_POST) { + telemetry.postDuration = duration; + } else if (stat_status == NIXL_TELEMETRY_POST_AND_FINISH) { + telemetry.postDuration = duration; + telemetry.xferDuration = duration; + } else { // stat_status == NIXL_TELEMETRY_FINISH + telemetry.xferDuration = duration; + } + + if (telemetry_pub && (stat_status != NIXL_TELEMETRY_POST)) { + telemetry_pub->addPostTime(telemetry.postDuration); + telemetry_pub->addXferTime(duration, backendOp == NIXL_WRITE, telemetry.totalBytes); + } + + NIXL_TRACE << "[NIXL TELEMETRY]: From backend " << engine->getType() + << nixl_post_status_str[stat_status] << " Xfer with " << telemetry.descCount + << " descriptors of total size " << telemetry.totalBytes << "B in " + << duration.count() << "us."; } /*** nixlAgentData constructor/destructor, as part of nixlAgent's ***/ -nixlAgentData::nixlAgentData(const std::string &name, - const nixlAgentConfig &cfg) : - name(name), config(cfg), lock(cfg.syncMode) -{ +nixlAgentData::nixlAgentData(const std::string &name, const nixlAgentConfig &cfg) + : name(name), + config(cfg), + lock(cfg.syncMode) { #if HAVE_ETCD if (getenv("NIXL_ETCD_ENDPOINTS")) { useEtcd = true; @@ -91,27 +118,51 @@ nixlAgentData::nixlAgentData(const std::string &name, useEtcd = false; NIXL_DEBUG << "NIXL ETCD is disabled"; } +#else + NIXL_DEBUG << "NIXL ETCD is excluded"; #endif // HAVE_ETCD if (name.empty()) throw std::invalid_argument("Agent needs a name"); memorySection = new nixlLocalSection(); + const char *telemetry_env_val = std::getenv(TELEMETRY_ENABLED_VAR); + const char *telemetry_env_dir = std::getenv(TELEMETRY_DIR_VAR); - const char *telemetry = std::getenv("NIXL_TELEMETRY_ENABLE"); - if (telemetry != nullptr) { - if (!strcasecmp(telemetry, "y")) + if (telemetry_env_val != nullptr) { + if (!strcasecmp(telemetry_env_val, "y") || !strcasecmp(telemetry_env_val, "1") || + !strcasecmp(telemetry_env_val, "yes") || !strcasecmp(telemetry_env_val, "on")) { telemetryEnabled = true; - else if (!strcasecmp(telemetry, "n")) - telemetryEnabled = false; - else + if (telemetry_env_dir != nullptr) { + std::string telemetry_file = std::string(telemetry_env_dir) + "/" + name; + telemetry_ = std::make_unique(telemetry_file, backendEngines); + NIXL_DEBUG << "NIXL telemetry is enabled with output file: " << telemetry_file; + } else { + NIXL_DEBUG << "NIXL telemetry is enabled without an output file"; + } + } else if (cfg.captureTelemetry) { + telemetryEnabled = true; + NIXL_WARN << "NIXL telemetry is enabled through config, " + "ignoring the NIXL_TELEMETRY_ENABLE environment variable"; + } else if (!strcasecmp(telemetry_env_val, "n") || !strcasecmp(telemetry_env_val, "0") || + !strcasecmp(telemetry_env_val, "no") || !strcasecmp(telemetry_env_val, "off")) { + NIXL_DEBUG << "NIXL telemetry is disabled"; + } else { NIXL_WARN - << "Invalid NIXL_TELEMETRY_ENABLE environment variable, not enabling telemetry."; + << "NIXL telemetry is disabled for invalid NIXL_TELEMETRY_ENABLE environment " + "variable -- valid are 'y', 'yes', '1', 'on', 'n', 'no', '0', 'off', any case"; + } + } else if (cfg.captureTelemetry) { + telemetryEnabled = true; + NIXL_DEBUG << "Capturing NIXL telemetry based on config (without an output file)"; } } nixlAgentData::~nixlAgentData() { delete memorySection; + // explicitly reset telemetry so i can publish backend events before destroying backends + telemetry_.reset(); + for (auto & elm: remoteSections) delete elm.second; @@ -205,6 +256,7 @@ nixlAgent::getPluginParams (const nixl_backend_t &type, return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "backend '" << type << "' not found"; return NIXL_ERR_NOT_FOUND; } @@ -212,8 +264,10 @@ nixl_status_t nixlAgent::getBackendParams (const nixlBackendH* backend, nixl_mem_list_t &mems, nixl_b_params_t ¶ms) const { - if (!backend) + if (!backend) { + NIXL_ERROR_FUNC << "backend handle is not provided"; return NIXL_ERR_INVALID_PARAM; + } NIXL_LOCK_GUARD(data->lock); mems = backend->engine->getSupportedMems(); @@ -235,8 +289,10 @@ nixlAgent::createBackend(const nixl_backend_t &type, NIXL_LOCK_GUARD(data->lock); // Registering same type of backend is not supported, unlikely and prob error - if (data->backendEngines.count(type)!=0) + if (data->backendEngines.count(type) != 0) { + NIXL_ERROR_FUNC << "backend already created for type '" << type << "'"; return NIXL_ERR_INVALID_PARAM; + } // Check if the plugin is in an illegal combination with another plugin backend already created for (const auto &combination : illegal_plugin_combinations) { @@ -244,20 +300,21 @@ nixlAgent::createBackend(const nixl_backend_t &type, for (const auto &plugin_name : combination) { if (plugin_name != type && data->backendEngines.find(plugin_name) != data->backendEngines.end()) { - NIXL_ERROR << "Plugin backend " << type << " is in illegal combination with " - << plugin_name; + NIXL_ERROR_FUNC << "Plugin backend " << type + << " is in illegal combination with " << plugin_name; return NIXL_ERR_NOT_ALLOWED; } } } } - init_params.localAgent = data->name; - init_params.type = type; - init_params.customParams = const_cast(¶ms); + init_params.localAgent = data->name; + init_params.type = type; + init_params.customParams = const_cast(¶ms); init_params.enableProgTh = data->config.useProgThread; - init_params.pthrDelay = data->config.pthrDelay; - init_params.syncMode = data->config.syncMode; + init_params.pthrDelay = data->config.pthrDelay; + init_params.syncMode = data->config.syncMode; + init_params.enableTelemetry_ = data->telemetry_ != nullptr; // First, try to load the backend as a plugin auto& plugin_manager = nixlPluginManager::getInstance(); @@ -267,20 +324,29 @@ nixlAgent::createBackend(const nixl_backend_t &type, // Plugin found, use it to create the backend backend = plugin_handle->createEngine(&init_params); } else { - NIXL_ERROR << "Unsupported backend: " << type; + NIXL_ERROR_FUNC << "unsupported backend '" << type << "'"; return NIXL_ERR_NOT_FOUND; } if (backend) { if (backend->getInitErr()) { delete backend; + NIXL_ERROR_FUNC << "backend initialization error for '" << type << "'"; return NIXL_ERR_BACKEND; } if (backend->supportsRemote()) { + if (!backend->supportsNotif()) { + delete backend; + NIXL_ERROR_FUNC << "backend '" << type << "' supportsRemote but not notifications"; + return NIXL_ERR_BACKEND; + } + ret = backend->getConnInfo(str); if (ret != NIXL_SUCCESS) { delete backend; + NIXL_ERROR_FUNC << "failed to get connection info for '" << type << "' with status " + << ret; return ret; } data->connMD[type] = str; @@ -291,6 +357,9 @@ nixlAgent::createBackend(const nixl_backend_t &type, if (NIXL_SUCCESS != ret) { delete backend; + NIXL_ERROR_FUNC + << "backend '" << type + << "' encountered error during intra-agent transfer setup with status " << ret; return ret; } } @@ -298,6 +367,7 @@ nixlAgent::createBackend(const nixl_backend_t &type, bknd_hndl = new nixlBackendH(backend); if (!bknd_hndl) { delete backend; + NIXL_ERROR_FUNC << "allocation of backend handle failed for '" << type << "'"; return NIXL_ERR_BACKEND; } @@ -322,6 +392,7 @@ nixlAgent::createBackend(const nixl_backend_t &type, return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "backend creation failed for '" << type << "'"; return NIXL_ERR_BACKEND; } @@ -331,6 +402,7 @@ nixlAgent::queryMem(const nixl_reg_dlist_t &descs, const nixl_opt_args_t *extra_params) const { if (!extra_params || extra_params->backends.size() != 1) { + NIXL_ERROR_FUNC << "this method requires exactly one backend to be passed"; return NIXL_ERR_INVALID_PARAM; } @@ -348,8 +420,10 @@ nixlAgent::registerMem(const nixl_reg_dlist_t &descs, NIXL_LOCK_GUARD(data->lock); if (!extra_params || extra_params->backends.size() == 0) { backend_list = &data->memToBackend[descs.getType()]; - if (backend_list->empty()) + if (backend_list->empty()) { + NIXL_ERROR_FUNC << "no available backends for mem type '" << descs.getType() << "'"; return NIXL_ERR_NOT_FOUND; + } } else { backend_list = new backend_list_t(); for (auto & elm : extra_params->backends) @@ -361,7 +435,7 @@ nixlAgent::registerMem(const nixl_reg_dlist_t &descs, for (size_t i=0; isize(); ++i) { nixlBackendEngine* backend = (*backend_list)[i]; // meta_descs use to be passed to loadLocalData - nixl_sec_dlist_t sec_descs(descs.getType(), false); + nixl_sec_dlist_t sec_descs(descs.getType()); ret = data->memorySection->addDescList(descs, backend, sec_descs); if (ret == NIXL_SUCCESS) { if (backend->supportsLocal()) { @@ -384,10 +458,20 @@ nixlAgent::registerMem(const nixl_reg_dlist_t &descs, if (extra_params && extra_params->backends.size() > 0) delete backend_list; - if (count > 0) + if (count > 0) { + // sum all the sizes of the descriptors using std::accumulate + if (data->telemetry_) { + uint64_t total_size = std::accumulate( + descs.begin(), + descs.end(), + uint64_t{0}, + [](uint64_t sum, const nixlBlobDesc &desc) { return sum + desc.len; }); + data->telemetry_->updateMemoryRegistered(total_size); + } return NIXL_SUCCESS; - else - return NIXL_ERR_BACKEND; + } + NIXL_ERROR_FUNC << "registration failed for the specified or all potential backends"; + return NIXL_ERR_BACKEND; } nixl_status_t @@ -403,8 +487,10 @@ nixlAgent::deregisterMem(const nixl_reg_dlist_t &descs, backend_set_t* avail_backends; avail_backends = data->memorySection->queryBackends( descs.getType()); - if (!avail_backends || avail_backends->empty()) + if (!avail_backends || avail_backends->empty()) { + NIXL_ERROR_FUNC << "no available backends for mem type '" << descs.getType() << "'"; return NIXL_ERR_NOT_FOUND; + } // Make a copy as we might change it in remDescList backend_set = *avail_backends; } else { @@ -418,7 +504,18 @@ nixlAgent::deregisterMem(const nixl_reg_dlist_t &descs, if (ret != NIXL_SUCCESS) bad_ret = ret; } - + if (bad_ret == NIXL_SUCCESS) { + if (data->telemetry_) { + uint64_t total_size = std::accumulate( + descs.begin(), + descs.end(), + uint64_t{0}, + [](uint64_t sum, const nixlBlobDesc &desc) { return sum + desc.len; }); + data->telemetry_->updateMemoryDeregistered(total_size); + } + } else { + NIXL_ERROR_FUNC << "deregistration failed on at least one backend with status " << bad_ret; + } return bad_ret; } @@ -431,12 +528,17 @@ nixlAgent::makeConnection(const std::string &remote_agent, int count = 0; NIXL_LOCK_GUARD(data->lock); - if (data->remoteBackends.count(remote_agent) == 0) + if (data->remoteBackends.count(remote_agent) == 0) { + NIXL_ERROR_FUNC << "metadata for remote agent '" << remote_agent << "' not found"; return NIXL_ERR_NOT_FOUND; + } if (!extra_params || extra_params->backends.size() == 0) { - if (data->remoteBackends[remote_agent].empty()) + if (data->remoteBackends[remote_agent].empty()) { + NIXL_ERROR_FUNC << "no backends are found in metadata for remote agent '" + << remote_agent << "'"; return NIXL_ERR_NOT_FOUND; + } for (auto & [r_bknd, conn_info] : data->remoteBackends[remote_agent]) backend_set.insert(r_bknd); } else { @@ -449,18 +551,24 @@ nixlAgent::makeConnection(const std::string &remote_agent, if (data->backendEngines.count(backend)!=0) { eng = data->backendEngines[backend]; ret = eng->connect(remote_agent); - if (ret) + if (ret) { + NIXL_ERROR_FUNC << "connect('" << remote_agent << "') failed on backend '" + << eng->getType() << "' with status " << ret; break; + } count++; } } - if (ret) + if (ret) // Error is already logged return ret; - else if (count == 0) // No common backend + + if (count == 0) { // No common backend + NIXL_ERROR_FUNC << "no common backend to connect with '" << remote_agent << "'"; return NIXL_ERR_BACKEND; - else - return NIXL_SUCCESS; + } + + return NIXL_SUCCESS; } nixl_status_t @@ -478,8 +586,11 @@ nixlAgent::prepXferDlist (const std::string &agent_name, NIXL_LOCK_GUARD(data->lock); // When central KV is supported, still it should return error, // just we can add a call to fetchRemoteMD for next time - if (!init_side && (data->remoteSections.count(agent_name) == 0)) + if (!init_side && (data->remoteSections.count(agent_name) == 0)) { + NIXL_ERROR_FUNC << "metadata for remote agent '" << agent_name << "' not found"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; + } if (!extra_params || extra_params->backends.size() == 0) { if (!init_side) @@ -489,8 +600,11 @@ nixlAgent::prepXferDlist (const std::string &agent_name, backend_set = data->memorySection-> queryBackends(descs.getType()); - if (!backend_set || backend_set->empty()) + if (!backend_set || backend_set->empty()) { + NIXL_ERROR_FUNC << "no available backends for mem type '" << descs.getType() << "'"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; + } } else { backend_set = new backend_set_t(); for (auto & elm : extra_params->backends) @@ -509,9 +623,7 @@ nixlAgent::prepXferDlist (const std::string &agent_name, } for (auto & backend : *backend_set) { - handle->descs[backend] = new nixl_meta_dlist_t ( - descs.getType(), - descs.isSorted()); + handle->descs[backend] = new nixl_meta_dlist_t(descs.getType()); if (init_side) ret = data->memorySection->populate( descs, backend, *(handle->descs[backend])); @@ -532,6 +644,10 @@ nixlAgent::prepXferDlist (const std::string &agent_name, if (count == 0) { delete handle; dlist_hndl = nullptr; + NIXL_ERROR_FUNC << "failed to prepare the descriptors for any of " + "the specified or potential backends for agent '" + << agent_name << "'"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } else { dlist_hndl = handle; @@ -555,16 +671,24 @@ nixlAgent::makeXferReq (const nixl_xfer_op_t &operation, req_hndl = nullptr; - if (!local_side || !remote_side) + if (!local_side || !remote_side) { + NIXL_ERROR_FUNC << "local or remote side handle is null"; + data->addErrorTelemetry(NIXL_ERR_INVALID_PARAM); return NIXL_ERR_INVALID_PARAM; + } - if ((!local_side->isLocal) || (remote_side->isLocal)) + if ((!local_side->isLocal) || (remote_side->isLocal)) { + NIXL_ERROR_FUNC << "invalid sides (local must be local, remote must be remote)"; + data->addErrorTelemetry(NIXL_ERR_INVALID_PARAM); return NIXL_ERR_INVALID_PARAM; + } NIXL_LOCK_GUARD(data->lock); // The remote was invalidated in between prepXferDlist and this call if (data->remoteSections.count(remote_side->remoteAgent) == 0) { - delete req_hndl; + NIXL_ERROR_FUNC << "remote agent '" << remote_side->remoteAgent + << "' was invalidated in between prepXferDlist and this call"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } @@ -589,28 +713,40 @@ nixlAgent::makeXferReq (const nixl_xfer_op_t &operation, } } - if (!backend) + if (!backend) { + NIXL_ERROR_FUNC << "could not find a common backend in the specified or " + "available list of backends for the prepped Dlists"; return NIXL_ERR_INVALID_PARAM; + } nixl_meta_dlist_t* local_descs = local_side->descs.at(backend); nixl_meta_dlist_t* remote_descs = remote_side->descs.at(backend); - size_t totalBytes = 0; + size_t total_bytes = 0; if ((desc_count == 0) || (remote_indices.size() == 0) || - (desc_count != (int) remote_indices.size())) + (desc_count != (int)remote_indices.size())) { + NIXL_ERROR_FUNC << "different number of indices for local (" << desc_count << "), remote (" + << remote_indices.size() << ")"; return NIXL_ERR_INVALID_PARAM; + } for (int i=0; i= local_descs->descCount()) - || (local_indices[i]<0)) + if ((local_indices[i] >= local_descs->descCount()) || (local_indices[i] < 0)) { + NIXL_ERROR_FUNC << "local index out of range at index " << i << " with value " + << local_indices[i]; return NIXL_ERR_INVALID_PARAM; - if ((remote_indices[i] >= remote_descs->descCount()) - || (remote_indices[i]<0)) + } + if ((remote_indices[i] >= remote_descs->descCount()) || (remote_indices[i] < 0)) { + NIXL_ERROR_FUNC << "remote index out of range at index " << i << " with value " + << remote_indices[i]; return NIXL_ERR_INVALID_PARAM; - if ((*local_descs )[local_indices [i]].len != - (*remote_descs)[remote_indices[i]].len) + } + if ((*local_descs)[local_indices[i]].len != (*remote_descs)[remote_indices[i]].len) { + NIXL_ERROR_FUNC << "length mismatch at index pair " << i << " with local index " + << local_indices[i] << " and remote index " << remote_indices[i]; return NIXL_ERR_INVALID_PARAM; - totalBytes += (*local_descs)[local_indices[i]].len; + } + total_bytes += (*local_descs)[local_indices[i]].len; } if (extra_params && extra_params->hasNotif) { @@ -619,19 +755,15 @@ nixlAgent::makeXferReq (const nixl_xfer_op_t &operation, } if ((opt_args.hasNotif) && (!backend->supportsNotif())) { + NIXL_ERROR_FUNC << "the selected backend '" << backend->getType() + << "' does not support notifications"; return NIXL_ERR_BACKEND; } - // Populate has been already done, no benefit in having sorted descriptors - // which will be overwritten by [] assignment operator. - nixlXferReqH* handle = new nixlXferReqH; - handle->initiatorDescs = new nixl_meta_dlist_t ( - local_descs->getType(), - false, desc_count); + std::unique_ptr handle = std::make_unique(); + handle->initiatorDescs = new nixl_meta_dlist_t(local_descs->getType(), desc_count); - handle->targetDescs = new nixl_meta_dlist_t ( - remote_descs->getType(), - false, desc_count); + handle->targetDescs = new nixl_meta_dlist_t(remote_descs->getType(), desc_count); if (extra_params && extra_params->skipDescMerge) { for (int i=0; ihasNotif = opt_args.hasNotif; handle->backendOp = operation; handle->status = NIXL_ERR_NOT_POSTED; - handle->telemetry.totalBytes = totalBytes; + + if (data->telemetryEnabled) { + handle->telemetry.totalBytes = total_bytes; + handle->telemetry.descCount = handle->initiatorDescs->descCount(); + } ret = handle->engine->prepXfer (handle->backendOp, *handle->initiatorDescs, @@ -693,11 +829,13 @@ nixlAgent::makeXferReq (const nixl_xfer_op_t &operation, handle->backendHandle, &opt_args); if (ret != NIXL_SUCCESS) { - delete handle; + NIXL_ERROR_FUNC << "backend '" << backend->getType() + << "' failed to prepare the transfer request with status " << ret; + data->addErrorTelemetry(ret); return ret; } - req_hndl = handle; + req_hndl = handle.release(); return NIXL_SUCCESS; } @@ -710,25 +848,32 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, const nixl_opt_args_t* extra_params) const { nixl_status_t ret1, ret2; nixl_opt_b_args_t opt_args; - backend_set_t* backend_set = new backend_set_t(); + + std::unique_ptr backend_set = std::make_unique(); req_hndl = nullptr; NIXL_SHARED_LOCK_GUARD(data->lock); if (data->remoteSections.count(remote_agent) == 0) { - delete backend_set; + NIXL_ERROR_FUNC << "metadata for remote agent '" << remote_agent << "' not found"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } + size_t total_bytes = 0; // Check the correspondence between descriptor lists - size_t totalBytes = 0; - if (local_descs.descCount() != remote_descs.descCount()) + if (local_descs.descCount() != remote_descs.descCount()) { + NIXL_ERROR_FUNC << "different descriptor list sizes (local=" << local_descs.descCount() + << ", remote=" << remote_descs.descCount() << ")"; return NIXL_ERR_INVALID_PARAM; + } for (int i = 0; i < local_descs.descCount(); ++i) { - if (local_descs[i].len != remote_descs[i].len) + if (local_descs[i].len != remote_descs[i].len) { + NIXL_ERROR_FUNC << "length mismatch at index " << i; return NIXL_ERR_INVALID_PARAM; - totalBytes += local_descs[i].len; + } + total_bytes += local_descs[i].len; } if (!extra_params || extra_params->backends.size() == 0) { @@ -740,7 +885,8 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, data->remoteSections[remote_agent]->queryBackends( remote_descs.getType()); if (!local_set || !remote_set) { - delete backend_set; + NIXL_ERROR_FUNC << "no backends found for local or remote for their " + "corresponding memory type"; return NIXL_ERR_NOT_FOUND; } @@ -749,7 +895,7 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, backend_set->insert(elm); if (backend_set->empty()) { - delete backend_set; + NIXL_ERROR_FUNC << "no potential backend found to be able to do the transfer"; return NIXL_ERR_NOT_FOUND; } } else { @@ -761,14 +907,10 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, // TODO: merge descriptors back to back in memory (like makeXferReq). // TODO [Perf]: Avoid heap allocation on the datapath, maybe use a mem pool - nixlXferReqH *handle = new nixlXferReqH; - handle->initiatorDescs = new nixl_meta_dlist_t ( - local_descs.getType(), - local_descs.isSorted()); + std::unique_ptr handle = std::make_unique(); + handle->initiatorDescs = new nixl_meta_dlist_t(local_descs.getType()); - handle->targetDescs = new nixl_meta_dlist_t ( - remote_descs.getType(), - remote_descs.isSorted()); + handle->targetDescs = new nixl_meta_dlist_t(remote_descs.getType()); // Currently we loop through and find first local match. Can use a // preference list or more exhaustive search. @@ -786,10 +928,10 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, } } - delete backend_set; - if (!handle->engine) { - delete handle; + NIXL_ERROR_FUNC << "no specified or potential backend had the required " + "registrations to be able to do the transfer"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } @@ -804,7 +946,9 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, } if (opt_args.hasNotif && (!handle->engine->supportsNotif())) { - delete handle; + NIXL_ERROR_FUNC << "the selected backend '" << handle->engine->getType() + << "' does not support notifications"; + data->addErrorTelemetry(NIXL_ERR_BACKEND); return NIXL_ERR_BACKEND; } @@ -813,7 +957,11 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, handle->status = NIXL_ERR_NOT_POSTED; handle->notifMsg = opt_args.notifMsg; handle->hasNotif = opt_args.hasNotif; - handle->telemetry.totalBytes = totalBytes; + + if (data->telemetryEnabled) { + handle->telemetry.totalBytes = total_bytes; + handle->telemetry.descCount = handle->initiatorDescs->descCount(); + } ret1 = handle->engine->prepXfer (handle->backendOp, *handle->initiatorDescs, @@ -822,11 +970,13 @@ nixlAgent::createXferReq(const nixl_xfer_op_t &operation, handle->backendHandle, &opt_args); if (ret1 != NIXL_SUCCESS) { - delete handle; + NIXL_ERROR_FUNC << "backend '" << handle->engine->getType() + << "' failed to prepare the transfer request with status " << ret1; + data->addErrorTelemetry(ret1); return ret1; } - req_hndl = handle; + req_hndl = handle.release(); return NIXL_SUCCESS; } @@ -837,52 +987,64 @@ nixlAgent::estimateXferCost(const nixlXferReqH *req_hndl, nixl_cost_t &method, const nixl_opt_args_t* extra_params) const { + nixl_status_t ret; NIXL_SHARED_LOCK_GUARD(data->lock); // Check if the remote agent connection info is still valid // (assuming cost estimation requires connection info like transfers) if (!req_hndl->remoteAgent.empty() && (data->remoteSections.count(req_hndl->remoteAgent) == 0)) { - NIXL_ERROR << "Invalid request handle: remote agent not found"; + NIXL_ERROR_FUNC << "invalid request handle, remote agent was invalidated " + "after transfer request creation"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } if (!req_hndl->engine) { - NIXL_ERROR << "Invalid request handle: engine is null"; + NIXL_ERROR_FUNC << "invalid request handle: engine is null"; + data->addErrorTelemetry(NIXL_ERR_UNKNOWN); return NIXL_ERR_UNKNOWN; } - return req_hndl->engine->estimateXferCost(req_hndl->backendOp, - *req_hndl->initiatorDescs, - *req_hndl->targetDescs, - req_hndl->remoteAgent, - req_hndl->backendHandle, - duration, - err_margin, - method, - extra_params); + ret = req_hndl->engine->estimateXferCost(req_hndl->backendOp, + *req_hndl->initiatorDescs, + *req_hndl->targetDescs, + req_hndl->remoteAgent, + req_hndl->backendHandle, + duration, + err_margin, + method, + extra_params); + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "backend '" << req_hndl->engine->getType() + << "' failed to estimate the transfer cost with status " << ret; + } + return ret; } nixl_status_t nixlAgent::postXferReq(nixlXferReqH *req_hndl, const nixl_opt_args_t* extra_params) const { - nixl_status_t ret; nixl_opt_b_args_t opt_args; opt_args.hasNotif = false; - if (!req_hndl) + if (!req_hndl) { + NIXL_ERROR_FUNC << "transfer request handle is null"; + data->addErrorTelemetry(NIXL_ERR_INVALID_PARAM); return NIXL_ERR_INVALID_PARAM; + } - // The initial checks should be fast if post succeeds, including them in the overall time if (data->telemetryEnabled) { - req_hndl->telemetry.startTime = std::chrono::high_resolution_clock::now(); + req_hndl->telemetry.startTime = std::chrono::steady_clock::now(); } NIXL_SHARED_LOCK_GUARD(data->lock); // Check if the remote was invalidated before post/repost if (data->remoteSections.count(req_hndl->remoteAgent) == 0) { - delete req_hndl; + NIXL_ERROR_FUNC << "remote agent '" << req_hndl->remoteAgent + << "' was invalidated after transfer request creation"; + data->addErrorTelemetry(NIXL_ERR_NOT_FOUND); return NIXL_ERR_NOT_FOUND; } @@ -891,9 +1053,16 @@ nixlAgent::postXferReq(nixlXferReqH *req_hndl, req_hndl->status = req_hndl->engine->checkXfer( req_hndl->backendHandle); if (req_hndl->status == NIXL_IN_PROG) { - delete req_hndl; + NIXL_ERROR_FUNC << "transfer request is still in progress and cannot be reposted"; return NIXL_ERR_REPOST_ACTIVE; } + + if (req_hndl->status == NIXL_ERR_REMOTE_DISCONNECT) { + data->invalidateRemoteData(req_hndl->remoteAgent); + NIXL_ERROR_FUNC << "remote agent '" << req_hndl->remoteAgent + << "' was disconnected after transfer request creation"; + return NIXL_ERR_REMOTE_DISCONNECT; + } } // Carrying over notification from xfer handle creation time @@ -916,51 +1085,99 @@ nixlAgent::postXferReq(nixlXferReqH *req_hndl, } if (opt_args.hasNotif && (!req_hndl->engine->supportsNotif())) { - delete req_hndl; + NIXL_ERROR_FUNC << "the selected backend '" << req_hndl->engine->getType() + << "' does not support notifications"; + data->addErrorTelemetry(NIXL_ERR_BACKEND); return NIXL_ERR_BACKEND; } // If status is not NIXL_IN_PROG we can repost, - ret = req_hndl->engine->postXfer (req_hndl->backendOp, - *req_hndl->initiatorDescs, - *req_hndl->targetDescs, - req_hndl->remoteAgent, - req_hndl->backendHandle, - &opt_args); - req_hndl->status = ret; + req_hndl->status = req_hndl->engine->postXfer(req_hndl->backendOp, + *req_hndl->initiatorDescs, + *req_hndl->targetDescs, + req_hndl->remoteAgent, + req_hndl->backendHandle, + &opt_args); + + if (req_hndl->status < 0) { + if (req_hndl->status == NIXL_ERR_REMOTE_DISCONNECT) { + NIXL_ERROR_FUNC << "remote agent '" << req_hndl->remoteAgent + << "' was disconnected after transfer request creation"; + data->invalidateRemoteData(req_hndl->remoteAgent); + return NIXL_ERR_REMOTE_DISCONNECT; + } else { + NIXL_ERROR_FUNC << "backend '" << req_hndl->engine->getType() + << "' failed to post the transfer request with status " + << req_hndl->status; + } + } if (data->telemetryEnabled) { - if (req_hndl->status == NIXL_SUCCESS) - req_hndl->updateRequestStats("Posted and Completed"); - else if (req_hndl->status == NIXL_IN_PROG) - req_hndl->updateRequestStats("Posted"); - // Errors should show up in debug log separately, not adding a print here + if (req_hndl->status < 0) { + data->addErrorTelemetry(req_hndl->status); + } else if (req_hndl->status == NIXL_IN_PROG) { + req_hndl->updateRequestStats(data->telemetry_, NIXL_TELEMETRY_POST); + } else { + req_hndl->updateRequestStats(data->telemetry_, NIXL_TELEMETRY_POST_AND_FINISH); + } } - return ret; + return req_hndl->status; } nixl_status_t nixlAgent::getXferStatus (nixlXferReqH *req_hndl) const { NIXL_SHARED_LOCK_GUARD(data->lock); - // If the status is done, no need to recheck. + // If the status is done, no need to recheck and no state changes. + // Same for users incorrectly recalling this method in error/done. if (req_hndl->status == NIXL_IN_PROG) { // Check if the remote was invalidated before completion if (data->remoteSections.count(req_hndl->remoteAgent) == 0) { - delete req_hndl; + NIXL_ERROR_FUNC << "remote agent '" << req_hndl->remoteAgent + << "' was invalidated during transfer"; return NIXL_ERR_NOT_FOUND; } - req_hndl->status = req_hndl->engine->checkXfer( - req_hndl->backendHandle); - } - if (data->telemetryEnabled && req_hndl->status == NIXL_SUCCESS) - req_hndl->updateRequestStats("Completed"); + req_hndl->status = req_hndl->engine->checkXfer(req_hndl->backendHandle); + if (req_hndl->status < 0) { + if (req_hndl->status == NIXL_ERR_REMOTE_DISCONNECT) { + data->invalidateRemoteData(req_hndl->remoteAgent); + return NIXL_ERR_REMOTE_DISCONNECT; + } else { + NIXL_ERROR_FUNC << "backend '" << req_hndl->engine->getType() + << "' returned error status " << req_hndl->status; + } + } + if (data->telemetryEnabled) { + if (req_hndl->status == NIXL_SUCCESS) { + req_hndl->updateRequestStats(data->telemetry_, NIXL_TELEMETRY_FINISH); + } else if (req_hndl->status < 0) { + data->addErrorTelemetry(req_hndl->status); + } + } + } + // If the status is error when entering this method, it was already logged return req_hndl->status; } +nixl_status_t +nixlAgent::getXferTelemetry(const nixlXferReqH *req_hndl, nixl_xfer_telem_t &telemetry) const { + + if (!data->telemetryEnabled) { + NIXL_ERROR_FUNC << "cannot return values when telemetry is not enabled."; + return NIXL_ERR_NO_TELEMETRY; + } + + if (req_hndl->status != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "Transfer is not complete yet"; + return req_hndl->status; + } + + telemetry = req_hndl->telemetry; + return NIXL_SUCCESS; +} nixl_status_t nixlAgent::queryXferBackend(const nixlXferReqH* req_hndl, @@ -984,9 +1201,12 @@ nixlAgent::releaseXferReq(nixlXferReqH *req_hndl) const { req_hndl->status = req_hndl->engine->releaseReqH( req_hndl->backendHandle); - if(req_hndl->status < 0) - return NIXL_ERR_REPOST_ACTIVE; - + if (req_hndl->status < 0) { + NIXL_ERROR_FUNC << "backend '" << req_hndl->engine->getType() + << "' could not release transfer request and returned error status " + << req_hndl->status; + return NIXL_ERR_REPOST_ACTIVE; // Might need renaming + } // just in case the backend doesn't set to NULL on success // this will prevent calling releaseReqH again in destructor req_hndl->backendHandle = nullptr; @@ -996,6 +1216,97 @@ nixlAgent::releaseXferReq(nixlXferReqH *req_hndl) const { return NIXL_SUCCESS; } +nixl_status_t +nixlAgent::createGpuXferReq(const nixlXferReqH &req_hndl, nixlGpuXferReqH &gpu_req_hndl) const { + if (!req_hndl.engine) { + NIXL_ERROR_FUNC << "Invalid request handle[" << &req_hndl << "]: engine is null"; + return NIXL_ERR_INVALID_PARAM; + } + + if (!req_hndl.backendHandle) { + NIXL_ERROR_FUNC << "Invalid request handle[" << &req_hndl << "]: backendHandle is null"; + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_SHARED_LOCK_GUARD(data->lock); + const auto status = req_hndl.engine->createGpuXferReq( + *req_hndl.backendHandle, *req_hndl.initiatorDescs, *req_hndl.targetDescs, gpu_req_hndl); + if (status == NIXL_SUCCESS) { + data->gpuReqToEngine.emplace(gpu_req_hndl, req_hndl.engine); + } + + return status; +} + +void +nixlAgent::releaseGpuXferReq(nixlGpuXferReqH gpu_req_hndl) const { + NIXL_SHARED_LOCK_GUARD(data->lock); + auto it = data->gpuReqToEngine.find(gpu_req_hndl); + if (it == data->gpuReqToEngine.end()) { + NIXL_WARN << "Invalid gpu_req_hndl[" << gpu_req_hndl << "] "; + return; + } + + it->second->releaseGpuXferReq(gpu_req_hndl); + + data->gpuReqToEngine.erase(it); +} + +nixl_status_t +nixlAgent::getGpuSignalSize(size_t &signal_size, const nixl_opt_args_t *extra_params) const { + if (!extra_params || extra_params->backends.empty()) { + NIXL_ERROR_FUNC << "backend must be specified in extra_params"; + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_SHARED_LOCK_GUARD(data->lock); + return extra_params->backends[0]->engine->getGpuSignalSize(signal_size); +} + +nixl_status_t +nixlAgent::prepGpuSignal(const nixl_reg_dlist_t &signal_descs, + const nixl_opt_args_t *extra_params) const { + if (signal_descs.descCount() == 0) { + NIXL_ERROR_FUNC << "signal descriptor list is empty"; + return NIXL_ERR_INVALID_PARAM; + } + + if (!extra_params || extra_params->backends.empty()) { + NIXL_ERROR_FUNC << "backend must be specified in extra_params"; + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_SHARED_LOCK_GUARD(data->lock); + + // Convert reg_dlist to xfer_dlist for populate call + nixl_xfer_dlist_t xfer_descs = signal_descs.trim(); + + nixlBackendH *backend = extra_params->backends[0]; + nixl_meta_dlist_t result(signal_descs.getType()); + nixl_status_t ret = data->memorySection->populate(xfer_descs, backend->engine, result); + + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "failed to populate signal metadata with specified backend"; + return ret; + } + + for (size_t i = 0; i < static_cast(result.descCount()); i++) { + void *signal = reinterpret_cast(result[i].addr); + ret = backend->engine->prepGpuSignal(*result[i].metadataP, signal); + + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "failed to prepare GPU signal " << i + << " with status: " << nixlEnumStrings::statusStr(ret); + return ret; + } + + NIXL_DEBUG << "Successfully prepared GPU signal " << i << " at address " << signal; + } + + NIXL_DEBUG << "Successfully prepared " << result.descCount() << " GPU signals"; + return NIXL_SUCCESS; +} + nixl_status_t nixlAgent::releasedDlistH (nixlDlistH* dlist_hndl) const { NIXL_LOCK_GUARD(data->lock); @@ -1013,8 +1324,10 @@ nixlAgent::getNotifs(nixl_notifs_t ¬if_map, NIXL_LOCK_GUARD(data->lock); if (!extra_params || extra_params->backends.size() == 0) { backend_list = &data->notifEngines; - if (backend_list->empty()) + if (backend_list->empty()) { + NIXL_ERROR_FUNC << "no backends support notifications"; return NIXL_ERR_BACKEND; + } } else { backend_list = new backend_list_t(); for (auto & elm : extra_params->backends) @@ -1022,6 +1335,7 @@ nixlAgent::getNotifs(nixl_notifs_t ¬if_map, backend_list->push_back(elm->engine); if (backend_list->empty()) { + NIXL_ERROR_FUNC << "none of specified backends support notifications"; delete backend_list; return NIXL_ERR_BACKEND; } @@ -1033,8 +1347,11 @@ nixlAgent::getNotifs(nixl_notifs_t ¬if_map, for (auto & eng: *backend_list) { bknd_notif_list.clear(); ret = eng->getNotifs(bknd_notif_list); - if (ret < 0) + if (ret < 0) { + NIXL_ERROR_FUNC << "backend '" << eng->getType() << "' returned error status " << ret + << " while getting notifications"; bad_ret=ret; + } if (bknd_notif_list.size() == 0) continue; @@ -1050,42 +1367,68 @@ nixlAgent::getNotifs(nixl_notifs_t ¬if_map, if (extra_params && extra_params->backends.size() > 0) delete backend_list; + // If any backend had an error, it was already logged return bad_ret; } nixl_status_t nixlAgent::genNotif(const std::string &remote_agent, const nixl_blob_t &msg, - const nixl_opt_args_t* extra_params) const { + const nixl_opt_args_t *extra_params) const { backend_list_t backend_list_value; - backend_list_t* backend_list; + backend_list_t *backend_list; + nixl_status_t ret; - NIXL_SHARED_LOCK_GUARD(data->lock); if (!extra_params || extra_params->backends.empty()) { backend_list = &data->notifEngines; - if (backend_list->empty()) - return NIXL_ERR_BACKEND; } else { backend_list = &backend_list_value; - for (auto &elm : extra_params->backends) - if (elm->engine->supportsNotif()) + for (auto &elm : extra_params->backends) { + if (elm->engine->supportsNotif()) { backend_list->push_back(elm->engine); + } + } + } - if (backend_list->empty()) { - return NIXL_ERR_BACKEND; + if (backend_list->empty()) { + NIXL_ERROR_FUNC << "no specified or potential backend supports notifications"; + return NIXL_ERR_BACKEND; + } + + NIXL_SHARED_LOCK_GUARD(data->lock); + + if (data->name == remote_agent) { + for (const auto &eng : *backend_list) { + if (eng->supportsLocal()) { + ret = eng->genNotif(remote_agent, msg); + if (ret < 0) { + NIXL_ERROR_FUNC << "backend '" << eng->getType() << "' returned error status " + << ret << " while sending intra-agent notifications"; + } + return ret; + } } + NIXL_ERROR_FUNC << "no specified or potential backend can send intra-agent notifications"; + return NIXL_ERR_NOT_FOUND; } + const auto iter = data->remoteBackends.find(remote_agent); - bool localNotif = data->name == remote_agent; - for (auto & eng: *backend_list) { - if ((localNotif && eng->supportsLocal()) || - (!localNotif && - data->remoteBackends[remote_agent].count(eng->getType()) != 0)) { - return eng->genNotif(remote_agent, msg); + if (iter != data->remoteBackends.end()) { + for (const auto &eng : *backend_list) { + if (iter->second.count(eng->getType()) != 0) { + ret = eng->genNotif(remote_agent, msg); + if (ret < 0) { + NIXL_ERROR_FUNC << "backend '" << eng->getType() << "' returned error status " + << ret << " while sending notification to agent '" + << remote_agent << "'"; + } + return ret; + } } } + NIXL_ERROR_FUNC << "no specified or potential backend could send the inter-agent notifications"; return NIXL_ERR_NOT_FOUND; } @@ -1099,35 +1442,36 @@ nixlAgent::getLocalMD (nixl_blob_t &str) const { // data->connMD was populated when the backend was created conn_cnt = data->connMD.size(); - if (conn_cnt == 0) // Error, no backend supports remote + if (conn_cnt == 0) { // Error, no backend supports remote + NIXL_ERROR_FUNC << "no backends support remote operations"; return NIXL_ERR_INVALID_PARAM; + } nixlSerDes sd; ret = sd.addStr("Agent", data->name); - if(ret) - return ret; + // Always returns SUCCESS, serdes class logs errors if necessary + if (ret) return NIXL_ERR_UNKNOWN; ret = sd.addBuf("Conns", &conn_cnt, sizeof(conn_cnt)); - if(ret) - return ret; + if (ret) return NIXL_ERR_UNKNOWN; for (auto &c : data->connMD) { nixl_backend = c.first; ret = sd.addStr("t", nixl_backend); - if(ret) - return ret; + if (ret) break; ret = sd.addStr("c", c.second); - if(ret) - return ret; + if (ret) break; } + if (ret) return NIXL_ERR_UNKNOWN; ret = sd.addStr("", "MemSection"); - if(ret) - return ret; + if (ret) return NIXL_ERR_UNKNOWN; ret = data->memorySection->serialize(&sd); - if(ret) + if (ret) { + NIXL_ERROR_FUNC << "serialization failed"; return ret; + } str = sd.exportStr(); return NIXL_SUCCESS; @@ -1147,8 +1491,10 @@ nixlAgent::getLocalPartialMD(const nixl_reg_dlist_t &descs, if (descs.descCount() != 0) { // Non-empty dlist, return backends that support the memory type backend_list = &data->memToBackend[descs.getType()]; - if (backend_list->empty()) + if (backend_list->empty()) { + NIXL_ERROR_FUNC << "no available backends for mem type '" << descs.getType() << "'"; return NIXL_ERR_NOT_FOUND; + } } else { // Empty dlist, return all backends backend_list = &tmp_list; @@ -1173,38 +1519,38 @@ nixlAgent::getLocalPartialMD(const nixl_reg_dlist_t &descs, selected_engines.insert(backend); } + if (selected_engines.size() == 0 && descs.descCount() > 0) { + NIXL_ERROR_FUNC << "no backends support the requested descriptors"; + return NIXL_ERR_BACKEND; + } + nixlSerDes sd; ret = sd.addStr("Agent", data->name); - if(ret) - return ret; + // Always returns SUCCESS, serdes class logs errors if necessary + if (ret) return NIXL_ERR_UNKNOWN; // Only add connection info if requested via extra_params or empty dlist size_t conn_cnt = ((extra_params && extra_params->includeConnInfo) || descs.descCount() == 0) ? found_iters.size() : 0; ret = sd.addBuf("Conns", &conn_cnt, sizeof(conn_cnt)); - if(ret) - return ret; + if (ret) return NIXL_ERR_UNKNOWN; for (size_t i = 0; i < conn_cnt; i++) { ret = sd.addStr("t", found_iters[i]->first); - if(ret) - return ret; + if (ret) break; ret = sd.addStr("c", found_iters[i]->second); - if(ret) - return ret; + if (ret) break; } - - // No engines found, but there are descs, this is an error - if (selected_engines.size() == 0 && descs.descCount() > 0) - return NIXL_ERR_BACKEND; + if (ret) return NIXL_ERR_UNKNOWN; ret = sd.addStr("", "MemSection"); - if(ret) - return ret; + if (ret) return NIXL_ERR_UNKNOWN; ret = data->memorySection->serializePartial(&sd, selected_engines, descs); - if(ret) + if (ret) { + NIXL_ERROR_FUNC << "serialization failed"; return ret; + } str = sd.exportStr(); return NIXL_SUCCESS; @@ -1213,88 +1559,73 @@ nixlAgent::getLocalPartialMD(const nixl_reg_dlist_t &descs, nixl_status_t nixlAgent::loadRemoteMD (const nixl_blob_t &remote_metadata, std::string &agent_name) { - int count = 0; nixlSerDes sd; - size_t conn_cnt; nixl_blob_t conn_info; nixl_backend_t nixl_backend; - nixlBackendEngine* eng; nixl_status_t ret; NIXL_LOCK_GUARD(data->lock); ret = sd.importStr(remote_metadata); - if(ret) - return ret; + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "failed to deserialize remote metadata"; + return NIXL_ERR_MISMATCH; + } std::string remote_agent = sd.getStr("Agent"); - if (remote_agent.size() == 0) + if (remote_agent.empty()) { + NIXL_ERROR_FUNC << "error in deserializing remote agent name"; return NIXL_ERR_MISMATCH; + } - if (remote_agent == data->name) + if (remote_agent == data->name) { + NIXL_ERROR_FUNC << "remote agent name same as local agent, " + "no need to load metadata"; return NIXL_ERR_INVALID_PARAM; + } NIXL_DEBUG << "Loading remote metadata for agent: " << remote_agent; + size_t conn_cnt; ret = sd.getBuf("Conns", &conn_cnt, sizeof(conn_cnt)); - if(ret) { - NIXL_ERROR << "Error getting connection count: " << nixlEnumStrings::statusStr(ret); - return ret; + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "error getting connection count: " << ret; + return NIXL_ERR_MISMATCH; } - for (size_t i=0; ibackendEngines.count(nixl_backend)!=0) { - - // No need to reload same conn info, error if it changed - if (data->remoteBackends.count(remote_agent) != 0 && - data->remoteBackends[remote_agent].count(nixl_backend) != 0) { - if (data->remoteBackends[remote_agent][nixl_backend] != conn_info) - return NIXL_ERR_NOT_ALLOWED; - count++; - continue; - } + if (nixl_backend.empty() || conn_info.empty()) { + NIXL_ERROR_FUNC << "failed to deserialize remote metadata"; + return NIXL_ERR_MISMATCH; + } - eng = data->backendEngines[nixl_backend]; - if (eng->supportsRemote()) { - ret = eng->loadRemoteConnInfo(remote_agent, conn_info); - if (ret) - return ret; // Error in load - count++; - data->remoteBackends[remote_agent].emplace(nixl_backend, conn_info); - } else { - // If there was an issue and we return error while some connections - // are loaded, they will be deleted in the backend destructor. - return NIXL_ERR_UNKNOWN; // This is an erroneous case - } + ret = data->loadConnInfo(remote_agent, nixl_backend, conn_info); + if (ret == NIXL_SUCCESS) { + count++; + } else if (ret != NIXL_ERR_NOT_SUPPORTED) { + NIXL_ERROR_FUNC << "error loading connection info for backend '" << nixl_backend + << "' with status " << ret; + return ret; } } - // No common backend, no point in loading the rest, unexpected - if (count == 0 && conn_cnt > 0) + if ((count == 0) && (conn_cnt > 0)) { + NIXL_ERROR_FUNC << "no common backend found"; return NIXL_ERR_BACKEND; + } - if (sd.getStr("") != "MemSection") + if (sd.getStr("") != "MemSection") { + NIXL_ERROR_FUNC << "failed to deserialize remote metadata"; return NIXL_ERR_MISMATCH; + } - if (data->remoteSections.count(remote_agent) == 0) - data->remoteSections[remote_agent] = new nixlRemoteSection( - remote_agent); - - ret = data->remoteSections[remote_agent]->loadRemoteData(&sd, - data->backendEngines); - - // TODO: can be more graceful, if just the new MD blob was improper - if (ret) { - delete data->remoteSections[remote_agent]; - data->remoteSections.erase(remote_agent); - data->remoteBackends.erase(remote_agent); + ret = data->loadRemoteSections(remote_agent, sd); + if (ret != NIXL_SUCCESS) { + NIXL_ERROR_FUNC << "error loading remote metadata for agent '" << remote_agent + << "' with status " << ret; return ret; } @@ -1306,23 +1637,30 @@ nixl_status_t nixlAgent::invalidateRemoteMD(const std::string &remote_agent) { NIXL_LOCK_GUARD(data->lock); - if (remote_agent == data->name) + if (remote_agent == data->name) { + NIXL_ERROR_FUNC << "remote agent same as local agent, cannot invalidate local metadata"; return NIXL_ERR_INVALID_PARAM; + } nixl_status_t ret = NIXL_ERR_NOT_FOUND; - if (data->remoteSections.count(remote_agent)!=0) { + if (data->remoteSections.count(remote_agent) != 0) { delete data->remoteSections[remote_agent]; data->remoteSections.erase(remote_agent); ret = NIXL_SUCCESS; } - if (data->remoteBackends.count(remote_agent)!=0) { - for (auto & it: data->remoteBackends[remote_agent]) + if (data->remoteBackends.count(remote_agent) != 0) { + for (auto &it : data->remoteBackends[remote_agent]) { data->backendEngines[it.first]->disconnect(remote_agent); + } + data->remoteBackends.erase(remote_agent); ret = NIXL_SUCCESS; } + if (ret != NIXL_SUCCESS) + NIXL_ERROR_FUNC << "error invalidating remote metadata for agent '" << remote_agent + << "' with status " << ret; return ret; } @@ -1330,7 +1668,10 @@ nixl_status_t nixlAgent::sendLocalMD (const nixl_opt_args_t* extra_params) const { nixl_blob_t myMD; nixl_status_t ret = getLocalMD(myMD); - if(ret < 0) return ret; + if (ret < 0) { + NIXL_ERROR_FUNC << "error getting local metadata with status " << ret; + return ret; + } // If IP is provided, use socket-based communication if (extra_params && !extra_params->ipAddr.empty()) { @@ -1344,8 +1685,11 @@ nixlAgent::sendLocalMD (const nixl_opt_args_t* extra_params) const { data->enqueueCommWork(std::make_tuple(ETCD_SEND, default_metadata_label, 0, std::move(myMD))); return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "invalid parameters to be used for either socket or ETCD"; return NIXL_ERR_INVALID_PARAM; #else + NIXL_ERROR_FUNC + << "sendLocalMD: ETCD is not supported and socket information was not provided either"; return NIXL_ERR_NOT_SUPPORTED; #endif // HAVE_ETCD } @@ -1355,7 +1699,10 @@ nixlAgent::sendLocalPartialMD(const nixl_reg_dlist_t &descs, const nixl_opt_args_t* extra_params) const { nixl_blob_t myMD; nixl_status_t ret = getLocalPartialMD(descs, myMD, extra_params); - if(ret < 0) return ret; + if (ret < 0) { + NIXL_ERROR_FUNC << "error getting local partial metadata with status " << ret; + return ret; + } // If IP is provided, use socket-based communication if (extra_params && !extra_params->ipAddr.empty()) { @@ -1367,14 +1714,16 @@ nixlAgent::sendLocalPartialMD(const nixl_reg_dlist_t &descs, // If no IP is provided, use etcd (now via thread) if (data->useEtcd) { if (!extra_params || extra_params->metadataLabel.empty()) { - NIXL_ERROR << "Metadata label is required for etcd send of local partial metadata"; + NIXL_ERROR_FUNC << "metadata label is required for etcd send of local partial metadata"; return NIXL_ERR_INVALID_PARAM; } data->enqueueCommWork(std::make_tuple(ETCD_SEND, extra_params->metadataLabel, 0, std::move(myMD))); return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "invalid parameters to be used for either socket or ETCD"; return NIXL_ERR_INVALID_PARAM; #else + NIXL_ERROR_FUNC << "ETCD is not supported and socket information was not provided either"; return NIXL_ERR_NOT_SUPPORTED; #endif // HAVE_ETCD } @@ -1397,8 +1746,10 @@ nixlAgent::fetchRemoteMD (const std::string remote_name, data->enqueueCommWork(std::make_tuple(ETCD_FETCH, std::move(metadata_label), 0, remote_name)); return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "invalid parameters to be used for either socket or ETCD"; return NIXL_ERR_INVALID_PARAM; #else + NIXL_ERROR_FUNC << "ETCD is not supported and socket information was not provided either"; return NIXL_ERR_NOT_SUPPORTED; #endif // HAVE_ETCD } @@ -1417,8 +1768,10 @@ nixlAgent::invalidateLocalMD (const nixl_opt_args_t* extra_params) const { data->enqueueCommWork(std::make_tuple(ETCD_INVAL, "", 0, "")); return NIXL_SUCCESS; } + NIXL_ERROR_FUNC << "invalid parameters to be used for either socket or ETCD"; return NIXL_ERR_INVALID_PARAM; #else + NIXL_ERROR_FUNC << "ETCD is not supported and socket information was not provided either"; return NIXL_ERR_NOT_SUPPORTED; #endif // HAVE_ETCD } @@ -1431,7 +1784,7 @@ nixlAgent::checkRemoteMD (const std::string remote_name, if (descs.descCount() == 0) { return NIXL_SUCCESS; } else { - nixl_meta_dlist_t dummy(descs.getType(), descs.isSorted()); + nixl_meta_dlist_t dummy(descs.getType()); // We only add to data->remoteBackends if data->backendEngines[backend] exists for (const auto& [backend, conn_info] : data->remoteBackends[remote_name]) if (data->remoteSections[remote_name]->populate( @@ -1440,5 +1793,7 @@ nixlAgent::checkRemoteMD (const std::string remote_name, dummy.clear(); } } + + // This is a checker method, returning not found is not an error to be logged return NIXL_ERR_NOT_FOUND; } diff --git a/src/core/nixl_listener.cpp b/src/core/nixl_listener.cpp index 5af242e67e..0dc5e76f35 100644 --- a/src/core/nixl_listener.cpp +++ b/src/core/nixl_listener.cpp @@ -27,6 +27,7 @@ #include #endif // HAVE_ETCD #include +#include const std::string default_metadata_label = "metadata"; @@ -59,26 +60,30 @@ int connectToIP(std::string ip_addr, int port) { return -1; } - // Use select to wait for connection with timeout - fd_set write_fds; - FD_ZERO(&write_fds); - FD_SET(ret_fd, &write_fds); + // Use poll to wait for connection with timeout + struct pollfd pfd; + pfd.fd = ret_fd; + pfd.events = POLLOUT; + pfd.revents = 0; - struct timeval tv; - tv.tv_sec = 1; - tv.tv_usec = 0; - - ret = select(ret_fd + 1, NULL, &write_fds, NULL, &tv); + ret = poll(&pfd, 1, 1000); // 1000ms timeout if (ret <= 0) { if (ret < 0) { - NIXL_PERROR << "select failed for ip_addr: " << ip_addr << " and port: " << port; + NIXL_PERROR << "poll failed for ip_addr: " << ip_addr << " and port: " << port; } else { - NIXL_ERROR << "select timed out for ip_addr: " << ip_addr << " and port: " << port; + NIXL_ERROR << "poll timed out for ip_addr: " << ip_addr << " and port: " << port; } close(ret_fd); return -1; } + if (!(pfd.revents & POLLOUT)) { + NIXL_ERROR << "poll returned but socket not ready for write for ip_addr: " << ip_addr + << " and port: " << port; + close(ret_fd); + return -1; + } + // Check if connection was successful int error = 0; socklen_t len = sizeof(error); @@ -182,6 +187,7 @@ class nixlEtcdClient { std::mutex invalidated_agents_mutex; std::unordered_map, std::hash, strEqual> agentWatchers; + std::chrono::microseconds watchTimeout_; // Helper function to create etcd key std::string makeKey(const std::string& agent_name, @@ -192,7 +198,9 @@ class nixlEtcdClient { } public: - nixlEtcdClient(const std::string& my_agent_name) { + nixlEtcdClient(const std::string &my_agent_name, + const std::chrono::microseconds &timeout = std::chrono::microseconds(5000000)) + : watchTimeout_(timeout) { const char* etcd_endpoints = std::getenv("NIXL_ETCD_ENDPOINTS"); if (!etcd_endpoints || strlen(etcd_endpoints) == 0) { throw std::runtime_error("No etcd endpoints provided"); @@ -314,9 +322,15 @@ class nixlEtcdClient { int64_t watch_index = response.index(); std::promise ret_prom; auto future = ret_prom.get_future(); + std::atomic promise_set{false}; // This lambda assumes lifetime only inside this method auto watcher_callback = [&](etcd::Response response) -> void { + if (promise_set.exchange(true)) { + NIXL_DEBUG << "Ignoring subsequent watch event for key: " << metadata_key; + return; + } + if (!response.is_ok()) { NIXL_ERROR << "Watch failed for key: " << metadata_key << " : " << response.error_message(); @@ -335,7 +349,7 @@ class nixlEtcdClient { auto watcher = etcd::Watcher(*etcd, metadata_key, watch_index, watcher_callback); - auto status = future.wait_for(std::chrono::seconds(5)); + auto status = future.wait_for(watchTimeout_); if (status == std::future_status::timeout) { NIXL_ERROR << "Watch timed out for key: " << metadata_key; return NIXL_ERR_BACKEND; @@ -427,7 +441,7 @@ void nixlAgentData::commWorker(nixlAgent* myAgent){ std::unique_ptr etcdClient = nullptr; // useEtcd is set in nixlAgent constructor and is true if NIXL_ETCD_ENDPOINTS is set if(useEtcd) { - etcdClient = std::make_unique(name); + etcdClient = std::make_unique(name, config.etcdWatchTimeout); } #endif // HAVE_ETCD @@ -447,7 +461,12 @@ void nixlAgentData::commWorker(nixlAgent* myAgent){ socklen_t client_addrlen = sizeof(client_address); if (getpeername(new_fd, (sockaddr*)&client_address, &client_addrlen) == 0) { char client_ip[INET_ADDRSTRLEN]; - inet_ntop(AF_INET, &client_address.sin_addr, client_ip, INET_ADDRSTRLEN); + if (inet_ntop(AF_INET, &client_address.sin_addr, client_ip, INET_ADDRSTRLEN) == + nullptr) { + NIXL_PERROR << "inet_ntop failed for client address"; + close(new_fd); + throw std::runtime_error("inet_ntop failed for client address"); + } accepted_client.first = std::string(client_ip); accepted_client.second = client_address.sin_port; } else { @@ -652,3 +671,83 @@ void nixlAgentData::getCommWork(std::vector &req_list){ req_list = std::move(commQueue); commQueue.clear(); } + +nixl_status_t +nixlAgentData::loadConnInfo(const std::string &remote_name, + const nixl_backend_t &backend, + const nixl_blob_t &conn_info) { + if (backendEngines.count(backend) == 0) { + NIXL_DEBUG << "Agent " << name << " does not support a remote backend: " << backend; + return NIXL_ERR_NOT_SUPPORTED; + } + + // No need to reload same conn info, error if it changed + if ((remoteBackends.count(remote_name) != 0) && + (remoteBackends[remote_name].count(backend) != 0)) { + if (remoteBackends[remote_name][backend] != conn_info) { + return NIXL_ERR_NOT_ALLOWED; + } + + return NIXL_SUCCESS; + } + + nixlBackendEngine *eng = backendEngines[backend]; + if (!eng->supportsRemote()) { + NIXL_DEBUG << backend << " does not support remote operations"; + return NIXL_ERR_NOT_SUPPORTED; + } + + const nixl_status_t ret = eng->loadRemoteConnInfo(remote_name, conn_info); + if (ret != NIXL_SUCCESS) { + return ret; + } + + remoteBackends[remote_name].emplace(backend, conn_info); + return NIXL_SUCCESS; +} + +nixl_status_t +nixlAgentData::loadRemoteSections(const std::string &remote_name, nixlSerDes &sd) { + if (remoteSections.count(remote_name) == 0) { + remoteSections[remote_name] = new nixlRemoteSection(remote_name); + } + + const nixl_status_t ret = remoteSections[remote_name]->loadRemoteData(&sd, backendEngines); + // TODO: can be more graceful, if just the new MD blob was improper + if (ret != NIXL_SUCCESS) { + delete remoteSections[remote_name]; + remoteSections.erase(remote_name); + remoteBackends.erase(remote_name); + return ret; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlAgentData::invalidateRemoteData(const std::string &remote_name) { + if (remote_name == name) { + NIXL_ERROR << "Agent " << name << " cannot invalidate itself"; + return NIXL_ERR_INVALID_PARAM; + } + + nixl_status_t ret = NIXL_ERR_NOT_FOUND; + auto it_section = remoteSections.find(remote_name); + if (it_section != remoteSections.end()) { + delete it_section->second; + remoteSections.erase(it_section); + ret = NIXL_SUCCESS; + } + + auto it_backends = remoteBackends.find(remote_name); + if (it_backends != remoteBackends.end()) { + for (auto &it : it_backends->second) { + backendEngines[it.first]->disconnect(remote_name); + } + + remoteBackends.erase(it_backends); + ret = NIXL_SUCCESS; + } + + return ret; +} diff --git a/src/core/nixl_plugin_manager.cpp b/src/core/nixl_plugin_manager.cpp index 95929d19dc..e42aaa9ba9 100644 --- a/src/core/nixl_plugin_manager.cpp +++ b/src/core/nixl_plugin_manager.cpp @@ -390,43 +390,50 @@ const std::vector& nixlPluginManager::getStaticPlugins() { return static_plugins_; } +#define NIXL_REGISTER_STATIC_PLUGIN(name) \ + extern nixlBackendPlugin *createStatic##name##Plugin(); \ + registerStaticPlugin(#name, createStatic##name##Plugin); + void nixlPluginManager::registerBuiltinPlugins() { +#ifdef STATIC_PLUGIN_LIBFABRIC + NIXL_REGISTER_STATIC_PLUGIN(LIBFABRIC) +#endif + #ifdef STATIC_PLUGIN_UCX - extern nixlBackendPlugin* createStaticUcxPlugin(); - registerStaticPlugin("UCX", createStaticUcxPlugin); -#endif //STATIC_PLUGIN_UCX + NIXL_REGISTER_STATIC_PLUGIN(UCX) +#endif #ifdef STATIC_PLUGIN_UCX_MO - extern nixlBackendPlugin* createStaticUcxMoPlugin(); - registerStaticPlugin("UCX_MO", createStaticUcxMoPlugin); -#endif // STATIC_PLUGIN_UCX_MO + NIXL_REGISTER_STATIC_PLUGIN(UCX_MO) +#endif #ifdef STATIC_PLUGIN_GDS #ifndef DISABLE_GDS_BACKEND - extern nixlBackendPlugin* createStaticGdsPlugin(); - registerStaticPlugin("GDS", createStaticGdsPlugin); -#endif // DISABLE_GDS_BACKEND -#endif // STATIC_PLUGIN_GDS + NIXL_REGISTER_STATIC_PLUGIN(GDS) +#endif +#endif + +#ifdef STATIC_PLUGIN_GDS_MT + NIXL_REGISTER_STATIC_PLUGIN(GDS_MT) +#endif #ifdef STATIC_PLUGIN_POSIX - extern nixlBackendPlugin* createStaticPosixPlugin(); - registerStaticPlugin("POSIX", createStaticPosixPlugin); -#endif // STATIC_PLUGIN_POSIX + NIXL_REGISTER_STATIC_PLUGIN(POSIX) +#endif #ifdef STATIC_PLUGIN_GPUNETIO -#ifndef DISABLE_GPUNETIO_BACKEND - extern nixlBackendPlugin* createStaticGpunetioPlugin(); - registerStaticPlugin("GPUNETIO", createStaticGpunetioPlugin); -#endif // DISABLE_GPUNETIO_BACKEND -#endif // STATIC_PLUGIN_GPUNETIO + NIXL_REGISTER_STATIC_PLUGIN(GPUNETIO) +#endif #ifdef STATIC_PLUGIN_OBJ - extern nixlBackendPlugin *createStaticObjPlugin(); - registerStaticPlugin ("OBJ", createStaticObjPlugin); -#endif // STATIC_PLUGIN_OBJ + NIXL_REGISTER_STATIC_PLUGIN(OBJ) +#endif #ifdef STATIC_PLUGIN_MOONCAKE - extern nixlBackendPlugin *createStaticMooncakePlugin(); - registerStaticPlugin("MOONCAKE", createStaticMooncakePlugin); -#endif // STATIC_PLUGIN_MOONCAKE + NIXL_REGISTER_STATIC_PLUGIN(MOONCAKE) +#endif + +#ifdef STATIC_PLUGIN_HF3FS + NIXL_REGISTER_STATIC_PLUGIN(HF3FS) +#endif } diff --git a/src/core/signalhandler.cpp b/src/core/signalhandler.cpp new file mode 100644 index 0000000000..40e6f3e409 --- /dev/null +++ b/src/core/signalhandler.cpp @@ -0,0 +1,88 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include +#include +#include +#include +#include +#include +#include + +#define MAX_BACKTRACE_DEPTH 100 + +void +print_backtrace() { + void *buffer[MAX_BACKTRACE_DEPTH]; + int nptrs = backtrace(buffer, MAX_BACKTRACE_DEPTH); + backtrace_symbols_fd(buffer, nptrs, STDERR_FILENO); +} + +void +gdb_signal_handler(int sig) { + signal(sig, SIG_DFL); + + const char *header = "\n!!! Caught signal. Generating backtrace: !!!\n"; + ssize_t ignored __attribute__((unused)) = write(STDERR_FILENO, header, strlen(header)); + + print_backtrace(); + + pid_t tid = fork(); + if (tid == 0) { + // Child process + char pid_buf[30] = {0}; + sprintf(pid_buf, "%d", getppid()); + + char exe_path_buf[1024]; + ssize_t len = readlink("/proc/self/exe", exe_path_buf, sizeof(exe_path_buf) - 1); + if (len != -1) { + exe_path_buf[len] = '\0'; + } else { + strcpy(exe_path_buf, "UNKNOWN_EXE"); + } + + // Replace child process with GDB + execlp("gdb", + "gdb", + "-q", + exe_path_buf, + pid_buf, + "--batch", + "-ex", + "thread apply all bt full", + "-ex", + "quit", + (char *)NULL); + + _exit(1); + } else if (tid > 0) { + // Parent process + int status; + waitpid(tid, &status, 0); + } + + // Re-raise signal to get core dump + raise(sig); +} + +__attribute__((constructor)) void +setup_gdb_handler() { + signal(SIGSEGV, gdb_signal_handler); + signal(SIGABRT, gdb_signal_handler); + signal(SIGFPE, gdb_signal_handler); + signal(SIGILL, gdb_signal_handler); +} diff --git a/src/core/telemetry.cpp b/src/core/telemetry.cpp new file mode 100644 index 0000000000..6c2ff6910e --- /dev/null +++ b/src/core/telemetry.cpp @@ -0,0 +1,253 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#include +#include +#include +#include +#include +#include +#include +#include + +#include "common/nixl_log.h" +#include "telemetry.h" +#include "telemetry_event.h" +#include "util.h" + +using namespace std::chrono_literals; +namespace fs = std::filesystem; + +constexpr std::chrono::milliseconds DEFAULT_TELEMETRY_RUN_INTERVAL = 100ms; +constexpr size_t DEFAULT_TELEMETRY_BUFFER_SIZE = 4096; + +nixlTelemetry::nixlTelemetry(const std::string &file_path, backend_map_t &backend_map) + : pool_(1), + writeTask_(pool_.get_executor(), DEFAULT_TELEMETRY_RUN_INTERVAL, false), + file_(file_path), + backendMap_(backend_map) { + if (file_path.empty()) { + throw std::invalid_argument("Telemetry file path cannot be empty"); + } + initializeTelemetry(); +} + +nixlTelemetry::~nixlTelemetry() { + writeTask_.enabled_ = false; + try { + writeTask_.timer_.cancel(); + pool_.stop(); + pool_.join(); + } + catch (const asio::system_error &e) { + NIXL_DEBUG << "Failed to cancel telemetry write timer: " << e.what(); + // continue anyway since it's not critical + } + + if (buffer_) { + writeEventHelper(); + buffer_.reset(); + } +} + +void +nixlTelemetry::initializeTelemetry() { + auto buffer_size = std::getenv(TELEMETRY_BUFFER_SIZE_VAR) ? + std::stoul(std::getenv(TELEMETRY_BUFFER_SIZE_VAR)) : + DEFAULT_TELEMETRY_BUFFER_SIZE; + + auto full_file_path = fs::path(file_); + + if (buffer_size == 0) { + throw std::invalid_argument("Telemetry buffer size cannot be 0"); + } + + NIXL_INFO << "Telemetry enabled, using buffer path: " << full_file_path + << " with size: " << buffer_size; + + buffer_ = std::make_unique>( + full_file_path, true, TELEMETRY_VERSION, buffer_size); + + auto run_interval = std::getenv(TELEMETRY_RUN_INTERVAL_VAR) ? + std::chrono::milliseconds(std::stoul(std::getenv(TELEMETRY_RUN_INTERVAL_VAR))) : + DEFAULT_TELEMETRY_RUN_INTERVAL; + + // Update write task interval and start it + writeTask_.callback_ = [this]() { return writeEventHelper(); }; + writeTask_.interval_ = run_interval; + writeTask_.enabled_ = true; + registerPeriodicTask(writeTask_); +} + +bool +nixlTelemetry::writeEventHelper() { + std::vector next_queue; + // assume next buffer will be the same size as the current one + next_queue.reserve(buffer_->capacity()); + { + std::lock_guard lock(mutex_); + events_.swap(next_queue); + } + for (auto &event : next_queue) { + // if full, ignore + buffer_->push(event); + } + // collect all events and sort them by timestamp + std::vector all_events; + for (auto &backend : backendMap_) { + auto backend_events = backend.second->getTelemetryEvents(); + for (auto &event : backend_events) { + // don't trust enum value coming from backend, + // as it might be different from the one in agent + event.category_ = nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND; + all_events.push_back(event); + } + } + std::sort(all_events.begin(), + all_events.end(), + [](const nixlTelemetryEvent &a, const nixlTelemetryEvent &b) { + return a.timestampUs_ < b.timestampUs_; + }); + for (auto &event : all_events) { + buffer_->push(event); + } + return true; +} + +void +nixlTelemetry::registerPeriodicTask(periodicTask &task) { + task.timer_.expires_after(task.interval_); + task.timer_.async_wait([this, &task](const asio::error_code &ec) { + if (ec != asio::error::operation_aborted) { + + task.callback_(); + + if (!task.enabled_) { + return; + } + + registerPeriodicTask(task); + } + }); +} + +void +nixlTelemetry::updateData(const std::string &event_name, + nixl_telemetry_category_t category, + uint64_t value) { + // agent can be multi-threaded + std::lock_guard lock(mutex_); + events_.emplace_back(std::chrono::duration_cast( + std::chrono::system_clock::now().time_since_epoch()) + .count(), + category, + event_name, + value); +} + +// The next 4 methods might be removed, as addXferTime covers them. +void +nixlTelemetry::updateTxBytes(uint64_t tx_bytes) { + updateData("agent_tx_bytes", nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, tx_bytes); +} + +void +nixlTelemetry::updateRxBytes(uint64_t rx_bytes) { + updateData("agent_rx_bytes", nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, rx_bytes); +} + +void +nixlTelemetry::updateTxRequestsNum(uint32_t tx_requests_num) { + updateData("agent_tx_requests_num", + nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, + tx_requests_num); +} + +void +nixlTelemetry::updateRxRequestsNum(uint32_t rx_requests_num) { + updateData("agent_rx_requests_num", + nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, + rx_requests_num); +} + +void +nixlTelemetry::updateErrorCount(nixl_status_t error_type) { + updateData( + nixlEnumStrings::statusStr(error_type), nixl_telemetry_category_t::NIXL_TELEMETRY_ERROR, 1); +} + +void +nixlTelemetry::updateMemoryRegistered(uint64_t memory_registered) { + updateData("agent_memory_registered", + nixl_telemetry_category_t::NIXL_TELEMETRY_MEMORY, + memory_registered); +} + +void +nixlTelemetry::updateMemoryDeregistered(uint64_t memory_deregistered) { + updateData("agent_memory_deregistered", + nixl_telemetry_category_t::NIXL_TELEMETRY_MEMORY, + memory_deregistered); +} + +void +nixlTelemetry::addXferTime(std::chrono::microseconds xfer_time, bool is_write, uint64_t bytes) { + std::string bytes_name; + std::string requests_name; + + if (is_write) { + bytes_name = "agent_tx_bytes"; + requests_name = "agent_tx_requests_num"; + } else { + bytes_name = "agent_rx_bytes"; + requests_name = "agent_rx_requests_num"; + } + auto time = std::chrono::duration_cast( + std::chrono::system_clock::now().time_since_epoch()) + .count(); + std::lock_guard lock(mutex_); + events_.emplace_back(time, + nixl_telemetry_category_t::NIXL_TELEMETRY_PERFORMANCE, + "agent_xfer_time", + xfer_time.count()); + events_.emplace_back( + time, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, bytes_name.c_str(), bytes); + events_.emplace_back( + time, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, requests_name.c_str(), 1); +} + +void +nixlTelemetry::addPostTime(std::chrono::microseconds post_time) { + updateData("agent_xfer_post_time", + nixl_telemetry_category_t::NIXL_TELEMETRY_PERFORMANCE, + post_time.count()); +} + +std::string +nixlEnumStrings::telemetryCategoryStr(const nixl_telemetry_category_t &category) { + static std::array nixl_telemetry_category_str = {"NIXL_TELEMETRY_MEMORY", + "NIXL_TELEMETRY_TRANSFER", + "NIXL_TELEMETRY_CONNECTION", + "NIXL_TELEMETRY_BACKEND", + "NIXL_TELEMETRY_ERROR", + "NIXL_TELEMETRY_PERFORMANCE", + "NIXL_TELEMETRY_SYSTEM", + "NIXL_TELEMETRY_CUSTOM", + "NIXL_TELEMETRY_MAX"}; + size_t category_int = static_cast(category); + if (category_int >= nixl_telemetry_category_str.size()) return "BAD_CATEGORY"; + return nixl_telemetry_category_str[category_int]; +} diff --git a/src/core/telemetry.h b/src/core/telemetry.h new file mode 100644 index 0000000000..2e9f1baffe --- /dev/null +++ b/src/core/telemetry.h @@ -0,0 +1,93 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef _TELEMETRY_H +#define _TELEMETRY_H + +#include "common/cyclic_buffer.h" +#include "telemetry_event.h" +#include "mem_section.h" +#include "nixl_types.h" + +#include +#include +#include +#include +#include +#include +#include + +#include + +struct periodicTask { + asio::steady_timer timer_; + std::function callback_; + std::chrono::milliseconds interval_; + std::atomic enabled_; + + periodicTask(const asio::any_io_executor &executor, + std::chrono::milliseconds interval, + bool enabled = false) + : timer_(executor), + callback_(nullptr), + interval_(interval), + enabled_(enabled) {} +}; + +class nixlTelemetry { +public: + nixlTelemetry(const std::string &file_path, backend_map_t &backend_map); + + ~nixlTelemetry(); + + void + updateTxBytes(uint64_t tx_bytes); + void + updateRxBytes(uint64_t rx_bytes); + void + updateTxRequestsNum(uint32_t num); + void + updateRxRequestsNum(uint32_t num); + void + updateErrorCount(nixl_status_t error_type); + void + updateMemoryRegistered(uint64_t memory_registered); + void + updateMemoryDeregistered(uint64_t memory_deregistered); + void + addXferTime(std::chrono::microseconds transaction_time, bool is_write, uint64_t bytes); + void + addPostTime(std::chrono::microseconds post_time); + +private: + void + initializeTelemetry(); + void + registerPeriodicTask(periodicTask &task); + void + updateData(const std::string &event_name, nixl_telemetry_category_t category, uint64_t value); + bool + writeEventHelper(); + std::unique_ptr> buffer_; + std::vector events_; + std::mutex mutex_; + asio::thread_pool pool_; + periodicTask writeTask_; + std::string file_; + backend_map_t &backendMap_; +}; + +#endif // _TELEMETRY_H diff --git a/src/core/telemetry_event.h b/src/core/telemetry_event.h new file mode 100644 index 0000000000..a37afa2a0d --- /dev/null +++ b/src/core/telemetry_event.h @@ -0,0 +1,74 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef _NIXL_TELEMETRY_H +#define _NIXL_TELEMETRY_H + +#include +#include + +#include "nixl_types.h" + +constexpr char TELEMETRY_BUFFER_SIZE_VAR[] = "NIXL_TELEMETRY_BUFFER_SIZE"; +constexpr char TELEMETRY_RUN_INTERVAL_VAR[] = "NIXL_TELEMETRY_RUN_INTERVAL"; + +constexpr int TELEMETRY_VERSION = 1; +constexpr size_t MAX_EVENT_NAME_LEN = 32; + +/** + * @enum nixl_telemetry_category_t + * @brief An enumeration of main telemetry event categories for easy filtering and aggregation + */ +enum class nixl_telemetry_category_t { + NIXL_TELEMETRY_MEMORY = 0, // Memory operations (register, deregister, allocation) + NIXL_TELEMETRY_TRANSFER = 1, // Data transfer operations (read, write) + NIXL_TELEMETRY_CONNECTION = 2, // Connection management (connect, disconnect) + NIXL_TELEMETRY_BACKEND = 3, // Backend-specific operations + NIXL_TELEMETRY_ERROR = 4, // Error events + NIXL_TELEMETRY_PERFORMANCE = 5, // Performance metrics + NIXL_TELEMETRY_SYSTEM = 6, // System-level events + NIXL_TELEMETRY_CUSTOM = 7, // Custom/user-defined events +}; + +namespace nixlEnumStrings { +std::string +telemetryCategoryStr(const nixl_telemetry_category_t &category); +} + +/** + * @struct nixlTelemetryEvent + * @brief A structure to hold individual telemetry event data for cyclic buffer storage + */ +struct nixlTelemetryEvent { + uint64_t timestampUs_; + nixl_telemetry_category_t category_; // Main event category for filtering + char eventName_[MAX_EVENT_NAME_LEN]; // Detailed event name/identifier + uint64_t value_; // Numeric value associated with the event + nixlTelemetryEvent() = default; + + nixlTelemetryEvent(uint64_t timestamp_us, + nixl_telemetry_category_t category, + const std::string &event_name, + uint64_t value) + : timestampUs_(timestamp_us), + category_(category), + value_(value) { + strncpy(eventName_, event_name.c_str(), MAX_EVENT_NAME_LEN - 1); + eventName_[MAX_EVENT_NAME_LEN - 1] = '\0'; + } +}; + +#endif // _NIXL_TELEMETRY_H diff --git a/src/core/transfer_request.h b/src/core/transfer_request.h index 8c0de0ed54..fb80b1abdb 100644 --- a/src/core/transfer_request.h +++ b/src/core/transfer_request.h @@ -17,8 +17,19 @@ #ifndef __TRANSFER_REQUEST_H_ #define __TRANSFER_REQUEST_H_ -constexpr auto min_chrono_time = std::chrono::time_point::min(); -using chrono_point_t = std::chrono::high_resolution_clock::time_point; +#include +#include +#include + +#include "nixl_types.h" +#include "backend_engine.h" +#include "telemetry.h" + +enum nixl_telemetry_stat_status_t { + NIXL_TELEMETRY_POST = 0, + NIXL_TELEMETRY_POST_AND_FINISH = 1, + NIXL_TELEMETRY_FINISH = 2 +}; // Contains pointers to corresponding backend engine and its handler, and populated // and verified DescLists, and other state and metadata needed for a NIXL transfer @@ -37,10 +48,7 @@ class nixlXferReqH { nixl_xfer_op_t backendOp; nixl_status_t status; - struct { - chrono_point_t startTime; - size_t totalBytes; - } telemetry; + nixl_xfer_telem_t telemetry; public: inline nixlXferReqH() { } @@ -54,7 +62,8 @@ class nixlXferReqH { } void - updateRequestStats(const std::string &dbg_msg_type); + updateRequestStats(std::unique_ptr &telemetry, + nixl_telemetry_stat_status_t stat_status); friend class nixlAgent; }; diff --git a/src/infra/mem_section.h b/src/infra/mem_section.h index f1bad541db..42ef27cbec 100644 --- a/src/infra/mem_section.h +++ b/src/infra/mem_section.h @@ -58,7 +58,37 @@ class nixlSectionDesc : public nixlMetaDesc { } }; -using nixl_sec_dlist_t = nixlDescList; +class nixlSecDescList : public nixlDescList { +public: + explicit nixlSecDescList(const nixl_mem_t &type) : nixlDescList(type, 0) {} + + using nixlDescList::operator[]; // bring in const overload + + void + addDesc(const nixlSectionDesc &desc) override; + + bool + verifySorted() const; + + nixlSectionDesc & + operator[](unsigned int index) override; + + int + getIndex(const nixlBasicDesc &query) const override; + + int + getCoveringIndex(const nixlBasicDesc &query) const; + + void + resize(const size_t &count) override; + + // Disable parent's convenience constructors that allow pre-sizing + nixlSecDescList(const nixlSecDescList &) = default; + nixlSecDescList & + operator=(const nixlSecDescList &) = default; +}; + +using nixl_sec_dlist_t = nixlSecDescList; using section_map_t = std::map; class nixlMemSection { diff --git a/src/infra/nixl_descriptors.cpp b/src/infra/nixl_descriptors.cpp index a020b2a60b..ecb0053891 100644 --- a/src/infra/nixl_descriptors.cpp +++ b/src/infra/nixl_descriptors.cpp @@ -146,13 +146,9 @@ void nixlBlobDesc::print(const std::string &suffix) const { // The template is used to select from nixlBasicDesc/nixlMetaDesc/nixlBlobDesc // There are no virtual functions, so the object is all data, no pointers. -template -nixlDescList::nixlDescList (const nixl_mem_t &type, - const bool &sorted, - const int &init_size) { +template nixlDescList::nixlDescList(const nixl_mem_t &type, const int &init_size) { static_assert (std::is_base_of::value); - this->type = type; - this->sorted = sorted; + this->type = type; this->descs.resize(init_size); } @@ -173,8 +169,7 @@ nixlDescList::nixlDescList(nixlSerDes* deserializer) { if (deserializer->getBuf("t", &type, sizeof(type))) return; - if (deserializer->getBuf("s", &sorted, sizeof(sorted))) - return; + if (deserializer->getBuf("n", &n_desc, sizeof(n_desc))) return; @@ -220,69 +215,13 @@ inline const T& nixlDescList::operator[](unsigned int index) const { // Setter template inline T& nixlDescList::operator[](unsigned int index) { - // To be added only in debug mode - // if (index >= descs.size()) - // throw std::out_of_range("Index is out of range"); - // sorted = false; + assert(index < descs.size()); return descs[index]; } template void nixlDescList::addDesc (const T &desc) { - if (!sorted) { - descs.push_back(desc); - } else { - // Since vector is kept soted, we can use upper_bound - auto itr = std::upper_bound(descs.begin(), descs.end(), desc); - if (itr == descs.end()) - descs.push_back(desc); - else - descs.insert(itr, desc); - } -} - -template -bool nixlDescList::overlaps (const T &desc, int &index) const { - if (!sorted) { - for (size_t i=0; i -bool nixlDescList::hasOverlaps () const { - if ((descs.size()==0) || (descs.size()==1)) - return false; - - if (!sorted) { - for (size_t i=0; i @@ -292,44 +231,19 @@ void nixlDescList::remDesc (const int &index){ descs.erase(descs.begin() + index); } -template -void nixlDescList::resize (const size_t &count) { - // To be added only in debug mode - // if (count > descs.size()) - // sorted = false; +template +void +nixlDescList::resize(const size_t &count) { descs.resize(count); } -template -bool nixlDescList::verifySorted() { - int size = (int) descs.size(); - if (size==0) { - return false; - } else if (size == 1) { - sorted = true; - return true; - } - - for (int i=0; i nixlDescList nixlDescList::trim() const { if constexpr (std::is_same::value) { return *this; } else { - nixlDescList trimmed(type, sorted); + nixlDescList trimmed(type); nixlBasicDesc* p; for (auto & elm: descs) { @@ -344,20 +258,9 @@ nixlDescList nixlDescList::trim() const { template int nixlDescList::getIndex(const nixlBasicDesc &query) const { - if (!sorted) { - auto itr = std::find(descs.begin(), descs.end(), query); - if (itr == descs.end()) - return NIXL_ERR_NOT_FOUND; // not found - return itr - descs.begin(); - } else { - auto itr = std::lower_bound(descs.begin(), descs.end(), query); - if (itr == descs.end()) - return NIXL_ERR_NOT_FOUND; // not found - // As desired, becomes nixlBasicDesc on both sides - if (*itr == query) - return itr - descs.begin(); - } - return NIXL_ERR_NOT_FOUND; + auto itr = std::find(descs.begin(), descs.end(), query); + if (itr == descs.end()) return NIXL_ERR_NOT_FOUND; // not found + return itr - descs.begin(); } template @@ -385,9 +288,6 @@ nixl_status_t nixlDescList::serialize(nixlSerDes* serializer) const { ret = serializer->addBuf("t", &type, sizeof(type)); if (ret) return ret; - ret = serializer->addBuf("s", &sorted, sizeof(sorted)); - if (ret) return ret; - ret = serializer->addBuf("n", &(n_desc), sizeof(n_desc)); if (ret) return ret; @@ -413,8 +313,7 @@ nixl_status_t nixlDescList::serialize(nixlSerDes* serializer) const { template void nixlDescList::print() const { - std::cout << "DescList of mem type " << type << " " - << (sorted ? "sorted" : "unsorted") << std::endl; + std::cout << "DescList of mem type " << type << std::endl; for (auto & elm : descs) { elm.print(""); } @@ -422,10 +321,7 @@ void nixlDescList::print() const { template bool operator==(const nixlDescList &lhs, const nixlDescList &rhs) { - if ((lhs.getType() != rhs.getType()) || - (lhs.descCount() != rhs.descCount()) || - (lhs.isSorted() != rhs.isSorted())) - return false; + if ((lhs.getType() != rhs.getType()) || (lhs.descCount() != rhs.descCount())) return false; for (size_t i=0; i(const nixlDescList &lhs, const nixlDescList &rhs); template bool operator==(const nixlDescList &lhs, const nixlDescList &rhs); + +// nixlSecDescList keeps the elements sorted +void +nixlSecDescList::addDesc(const nixlSectionDesc &desc) { + auto &vec = this->descs; + auto itr = std::upper_bound(vec.begin(), vec.end(), desc); + if (itr == vec.end()) + vec.push_back(desc); + else + vec.insert(itr, desc); +} + +bool +nixlSecDescList::verifySorted() const { + const auto &vec = this->descs; + int size = (int)vec.size(); + if (size <= 1) return (size == 1); + for (int i = 0; i < size - 1; ++i) { + if (vec[i + 1] < vec[i]) return false; + } + return true; +} + +nixlSectionDesc & +nixlSecDescList::operator[](unsigned int index) { + nixlSectionDesc &ref = this->descs[index]; + assert(verifySorted()); + return ref; +} + +int +nixlSecDescList::getIndex(const nixlBasicDesc &query) const { + auto itr = std::lower_bound(this->descs.begin(), this->descs.end(), query); + if (itr == this->descs.end()) return NIXL_ERR_NOT_FOUND; + if (static_cast(*itr) == query) + return static_cast(itr - this->descs.begin()); + return NIXL_ERR_NOT_FOUND; +} + +int +nixlSecDescList::getCoveringIndex(const nixlBasicDesc &query) const { + auto itr = std::lower_bound(this->descs.begin(), this->descs.end(), query); + if (itr != this->descs.end() && itr->covers(query)) + return static_cast(itr - this->descs.begin()); + // If query and element don't have the same start address, try previous entry + if (itr != this->descs.begin()) { + auto prev_itr = std::prev(itr, 1); + if (prev_itr->covers(query)) return static_cast(prev_itr - this->descs.begin()); + } + return -1; +} + +void +nixlSecDescList::resize(const size_t &count) { + if (count > this->descs.size()) + throw std::logic_error( + "nixlSecDescList: to keep list sorted, resize growth is not allowed."); + this->descs.resize(count); +} diff --git a/src/infra/nixl_memory_section.cpp b/src/infra/nixl_memory_section.cpp index e73006443f..d5cbed8a1a 100644 --- a/src/infra/nixl_memory_section.cpp +++ b/src/infra/nixl_memory_section.cpp @@ -15,6 +15,7 @@ * limitations under the License. */ #include +#include #include #include "nixl.h" #include "nixl_descriptors.h" @@ -39,10 +40,7 @@ nixl_status_t nixlMemSection::populate (const nixl_xfer_dlist_t &query, nixlBackendEngine* backend, nixl_meta_dlist_t &resp) const { - if (query.getType() != resp.getType()) - return NIXL_ERR_INVALID_PARAM; - // 1-to-1 mapping cannot hold - if (query.isSorted() != resp.isSorted()) + if ((query.getType() != resp.getType()) || (query.descCount() == 0)) return NIXL_ERR_INVALID_PARAM; section_key_t sec_key = std::make_pair(query.getType(), backend); @@ -50,98 +48,43 @@ nixl_status_t nixlMemSection::populate (const nixl_xfer_dlist_t &query, if (it==sectionMap.end()) return NIXL_ERR_NOT_FOUND; - nixlBasicDesc *p; nixl_sec_dlist_t* base = it->second; resp.resize(query.descCount()); - if (!base->isSorted()) { - int count = 0; - for (int i=0; idescCount(); - s_index = 0; - q_index = 0; - - while (q_indexcovers(*q)) { - p = &resp[q_index]; - *p = *q; - resp[q_index].metadataP = s->metadataP; - q_index++; - } else { - s_index++; - // TODO: add early termination if already (*q < *s), - // but s was not properly covering q - if (s_index==size) { - resp.clear(); - return NIXL_ERR_UNKNOWN; - } - } - } - - // To be added only in debug mode - // resp.verifySorted(); - return NIXL_SUCCESS; + int size = base->descCount(); + int s_index = 0; + // Use logN search for the first element, instead of linear search + s_index = base->getCoveringIndex(query[0]); + if (s_index < 0) { + resp.clear(); + return NIXL_ERR_UNKNOWN; + } + static_cast(resp[0]) = query[0]; + resp[0].metadataP = (*base)[s_index].metadataP; + + // Walk forward for non-decreasing elements; logN search on temporal disorder + for (int i = 1; i < query.descCount(); ++i) { + if (query[i] < query[i - 1]) { + // Disorder in the list, resolve this element using logN search + s_index = base->getCoveringIndex(query[i]); + if (s_index < 0) { + resp.clear(); + return NIXL_ERR_UNKNOWN; + } } else { - int last_found = 0; - for (int i=0; ibegin() + last_found, - base->end(), *q); - - // Same start address case - if (itr != base->end()){ - if (itr->covers(*q)) { - found = true; - } - } - - // query starts starts later, try previous entry - if ((!found) && (itr != base->begin())){ - itr = std::prev(itr , 1); - if (itr->covers(*q)) { - found = true; - } - } - - if (found) { - p = &resp[i]; - *p = *q; - resp[i].metadataP = itr->metadataP; - } else { - resp.clear(); - return NIXL_ERR_UNKNOWN; - } + while (s_index < size && !(*base)[s_index].covers(query[i])) + ++s_index; + if (s_index == size) { + resp.clear(); + return NIXL_ERR_UNKNOWN; } - return NIXL_SUCCESS; } + + static_cast(resp[i]) = query[i]; + resp[i].metadataP = (*base)[s_index].metadataP; } + return NIXL_SUCCESS; } /*** Class nixlLocalSection implementation ***/ @@ -159,7 +102,7 @@ nixl_status_t nixlLocalSection::addDescList (const nixl_reg_dlist_t &mem_elms, auto it = sectionMap.find(sec_key); if (it==sectionMap.end()) { // New desc list - sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem, true); + sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem); memToBackend[nixl_mem].insert(backend); } nixl_sec_dlist_t *target = sectionMap[sec_key]; @@ -319,7 +262,7 @@ nixl_status_t nixlLocalSection::serializePartial(nixlSerDes* serializer, // TODO: consider section_map_t to be a map of unique_ptr or instance of nixl_meta_dlist_t. // This will avoid the need to delete the nixl_sec_dlist_t instances. const nixl_sec_dlist_t *base = it->second; - nixl_sec_dlist_t *resp = new nixl_sec_dlist_t(nixl_mem, mem_elms.isSorted()); + nixl_sec_dlist_t *resp = new nixl_sec_dlist_t(nixl_mem); for (const auto &desc : mem_elms) { int index = base->getIndex(desc); if (index < 0) { @@ -368,8 +311,9 @@ nixl_status_t nixlRemoteSection::addDescList ( // Without it, its corrupt data, we keep the last option without raising an error nixl_mem_t nixl_mem = mem_elms.getType(); section_key_t sec_key = std::make_pair(nixl_mem, backend); - if (sectionMap.count(sec_key) == 0) - sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem, true); + if (sectionMap.count(sec_key) == 0) { + sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem); + } memToBackend[nixl_mem].insert(backend); // Fine to overwrite, it's a set nixl_sec_dlist_t *target = sectionMap[sec_key]; @@ -436,8 +380,9 @@ nixl_status_t nixlRemoteSection::loadLocalData ( nixl_mem_t nixl_mem = mem_elms.getType(); section_key_t sec_key = std::make_pair(nixl_mem, backend); - if (sectionMap.count(sec_key) == 0) - sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem, true); + if (sectionMap.count(sec_key) == 0) { + sectionMap[sec_key] = new nixl_sec_dlist_t(nixl_mem); + } memToBackend[nixl_mem].insert(backend); // Fine to overwrite, it's a set nixl_sec_dlist_t *target = sectionMap[sec_key]; diff --git a/src/plugins/cuda_gds/gds_backend.h b/src/plugins/cuda_gds/gds_backend.h index 2274c3d76a..428d414417 100644 --- a/src/plugins/cuda_gds/gds_backend.h +++ b/src/plugins/cuda_gds/gds_backend.h @@ -121,9 +121,6 @@ class nixlGdsEngine : public nixlBackendEngine { bool supportsLocal() const { return true; } - bool supportsProgTh() const { - return false; - } nixl_mem_list_t getSupportedMems() const { nixl_mem_list_t mems; diff --git a/src/plugins/cuda_gds/gds_plugin.cpp b/src/plugins/cuda_gds/gds_plugin.cpp index 833882338a..b8f8925b69 100644 --- a/src/plugins/cuda_gds/gds_plugin.cpp +++ b/src/plugins/cuda_gds/gds_plugin.cpp @@ -18,71 +18,23 @@ #include "backend/backend_plugin.h" #include "gds_backend.h" -// Plugin version information -static const char* PLUGIN_NAME = "GDS"; -static const char* PLUGIN_VERSION = "0.1.1"; -// Function to create a new GDS backend engine instance -static nixlBackendEngine* create_gds_engine(const nixlBackendInitParams* init_params) { - return new nixlGdsEngine(init_params); -} - -static void destroy_gds_engine(nixlBackendEngine* engine) { - delete engine; -} - -// Function to get the plugin name -static const char* get_plugin_name() { - return PLUGIN_NAME; -} - -// Function to get the plugin version -static const char* get_plugin_version() { - return PLUGIN_VERSION; -} - -// Function to get backend options -static nixl_b_params_t get_backend_options() { - nixl_b_params_t params; - return params; -} - -// Function to get supported backend mem types -static nixl_mem_list_t get_backend_mems() { - nixl_mem_list_t mems; - mems.push_back(DRAM_SEG); - mems.push_back(VRAM_SEG); - mems.push_back(FILE_SEG); - return mems; -} - -// Static plugin structure -static nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_gds_engine, - destroy_gds_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems -}; +// Plugin type alias for convenience +using gds_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_GDS - -nixlBackendPlugin* createStaticGdsPlugin() { - return &plugin; // Return the static plugin instance +nixlBackendPlugin * +createStaticGDSPlugin() { + return gds_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GDS", "0.1.1", {}, {DRAM_SEG, VRAM_SEG, FILE_SEG}); } - #else - -// Plugin initialization function -extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - return &plugin; -} - -// Plugin cleanup function -extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { - // Cleanup any resources if needed +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return gds_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GDS", "0.1.1", {}, {DRAM_SEG, VRAM_SEG, FILE_SEG}); } +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} #endif diff --git a/src/plugins/cuda_gds/gds_utils.cpp b/src/plugins/cuda_gds/gds_utils.cpp index 560fee2b23..b58d8bf153 100644 --- a/src/plugins/cuda_gds/gds_utils.cpp +++ b/src/plugins/cuda_gds/gds_utils.cpp @@ -23,7 +23,7 @@ nixl_status_t gdsUtil::registerFileHandle(int fd, gdsFileHandle& gds_handle) { CUfileError_t status; - CUfileDescr_t descr; + CUfileDescr_t descr = {}; CUfileHandle_t handle; descr.handle.fd = fd; diff --git a/src/plugins/cuda_gds/meson.build b/src/plugins/cuda_gds/meson.build index 8017cf0784..ea4516c9ec 100644 --- a/src/plugins/cuda_gds/meson.build +++ b/src/plugins/cuda_gds/meson.build @@ -46,7 +46,8 @@ else install: true, cpp_args: ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', 'echo "GDS=' + gds_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', diff --git a/src/plugins/gds_mt/gds_mt_backend.h b/src/plugins/gds_mt/gds_mt_backend.h index 0b9e203266..773a73b64b 100644 --- a/src/plugins/gds_mt/gds_mt_backend.h +++ b/src/plugins/gds_mt/gds_mt_backend.h @@ -52,10 +52,6 @@ class nixlGdsMtEngine : public nixlBackendEngine { supportsLocal() const override { return true; } - bool - supportsProgTh() const override { - return false; - } nixl_mem_list_t getSupportedMems() const override { diff --git a/src/plugins/gds_mt/gds_mt_plugin.cpp b/src/plugins/gds_mt/gds_mt_plugin.cpp index a8f7f910f1..63bab530b3 100644 --- a/src/plugins/gds_mt/gds_mt_plugin.cpp +++ b/src/plugins/gds_mt/gds_mt_plugin.cpp @@ -20,68 +20,23 @@ #include "common/nixl_log.h" #include -static const char *PLUGIN_NAME = "GDS_MT"; -static const char *PLUGIN_VERSION = "0.1.0"; -static nixlBackendEngine * -create_gds_mt_engine (const nixlBackendInitParams *init_params) { - try { - return new nixlGdsMtEngine (init_params); - } - catch (const std::exception &e) { - NIXL_ERROR << "GDS_MT: Failed to create engine: " << e.what(); - return nullptr; - } -} - -static void -destroy_gds_mt_engine (nixlBackendEngine *engine) { - delete engine; -} - -static const char * -get_plugin_name() { - return PLUGIN_NAME; -} - -static const char * -get_plugin_version() { - return PLUGIN_VERSION; -} - -static nixl_b_params_t -get_backend_options() { - return {}; -} - -static nixl_mem_list_t -get_backend_mems() { - return {DRAM_SEG, VRAM_SEG, FILE_SEG}; -} - -static nixlBackendPlugin plugin = {NIXL_PLUGIN_API_VERSION, - create_gds_mt_engine, - destroy_gds_mt_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems}; +// Plugin type alias for convenience +using gds_mt_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_GDS_MT - nixlBackendPlugin * -createStaticGdsMtPlugin() { - return &plugin; +createStaticGDS_MTPlugin() { + return gds_mt_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GDS_MT", "0.1.0", {}, {DRAM_SEG, VRAM_SEG, FILE_SEG}); } - #else - extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * nixl_plugin_init() { - return &plugin; + return gds_mt_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GDS_MT", "0.1.0", {}, {DRAM_SEG, VRAM_SEG, FILE_SEG}); } extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() {} - #endif diff --git a/src/plugins/gds_mt/meson.build b/src/plugins/gds_mt/meson.build index c30c97eaa1..fcc94100d8 100644 --- a/src/plugins/gds_mt/meson.build +++ b/src/plugins/gds_mt/meson.build @@ -46,7 +46,8 @@ else install: true, cpp_args: ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', 'echo "GDS_MT=' + gds_mt_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', diff --git a/src/plugins/gpunetio/gpunetio_backend.h b/src/plugins/gpunetio/gpunetio_backend.h index aed7361878..8f3a35117e 100644 --- a/src/plugins/gpunetio/gpunetio_backend.h +++ b/src/plugins/gpunetio/gpunetio_backend.h @@ -46,9 +46,6 @@ class nixlDocaEngine : public nixlBackendEngine { bool supportsNotif() const { return true; } - bool supportsProgTh() const { - return false; - } nixl_mem_list_t getSupportedMems() const; diff --git a/src/plugins/gpunetio/gpunetio_backend_aux.h b/src/plugins/gpunetio/gpunetio_backend_aux.h index cd6eaba130..ebbd137d2d 100644 --- a/src/plugins/gpunetio/gpunetio_backend_aux.h +++ b/src/plugins/gpunetio/gpunetio_backend_aux.h @@ -44,7 +44,6 @@ #include "nixl.h" // Local includes -#include "common/list_elem.h" #include "common/nixl_time.h" constexpr uint32_t DOCA_MAX_COMPLETION_INFLIGHT = 128; diff --git a/src/plugins/gpunetio/gpunetio_plugin.cpp b/src/plugins/gpunetio/gpunetio_plugin.cpp index 2970dddaee..5f6e68cc93 100644 --- a/src/plugins/gpunetio/gpunetio_plugin.cpp +++ b/src/plugins/gpunetio/gpunetio_plugin.cpp @@ -18,40 +18,9 @@ #include "backend/backend_plugin.h" #include "gpunetio_backend.h" -// Plugin version information -static const char *PLUGIN_NAME = "GPUNETIO"; -static const char *PLUGIN_VERSION = "0.1.0"; - -static nixlBackendEngine * -create_engine (const nixlBackendInitParams *init_params) { - try { - return new nixlDocaEngine (init_params); - } - catch (const std::exception &e) { - return nullptr; - } -} - -static void -destroy_engine (nixlBackendEngine *engine) { - delete engine; -} - -// Function to get the plugin name -static const char * -get_plugin_name() { - return PLUGIN_NAME; -} - -// Function to get the plugin version -static const char * -get_plugin_version() { - return PLUGIN_VERSION; -} - -// Function to get backend options -static nixl_b_params_t -get_backend_options() { +namespace { +nixl_b_params_t +get_gpunetio_options() { nixl_b_params_t params; params["network_devices"] = ""; params["gpu_devices"] = ""; @@ -59,42 +28,25 @@ get_backend_options() { return params; } -// Function to get supported backend mem types -static nixl_mem_list_t -get_backend_mems() { - nixl_mem_list_t mems; - mems.push_back (DRAM_SEG); - mems.push_back (VRAM_SEG); - return mems; -} -// Static plugin structure -static nixlBackendPlugin plugin = {NIXL_PLUGIN_API_VERSION, - create_engine, - destroy_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems}; +} // namespace -#ifdef STATIC_PLUGIN_GPUNETIO +// Plugin type alias for convenience +using gpunetio_plugin_t = nixlBackendPluginCreator; +#ifdef STATIC_PLUGIN_GPUNETIO nixlBackendPlugin * -createStaticDocaPlugin() { - return &plugin; // Return the static plugin instance +createStaticGPUNETIOPlugin() { + return gpunetio_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GPUNETIO", "0.1.0", get_gpunetio_options(), {DRAM_SEG, VRAM_SEG}); } - #else - -// Plugin initialization function extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * nixl_plugin_init() { - return &plugin; + return gpunetio_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "GPUNETIO", "0.1.0", get_gpunetio_options(), {DRAM_SEG, VRAM_SEG}); } -// Plugin cleanup function extern "C" NIXL_PLUGIN_EXPORT void -nixl_plugin_fini() { - // Cleanup any resources if needed -} +nixl_plugin_fini() {} #endif diff --git a/src/plugins/gpunetio/meson.build b/src/plugins/gpunetio/meson.build index 7df03adc71..f1331b9cc4 100644 --- a/src/plugins/gpunetio/meson.build +++ b/src/plugins/gpunetio/meson.build @@ -44,7 +44,8 @@ else cpp_args : compile_flags + ['-fPIC'], cuda_args : gpu_cuda_args, name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', diff --git a/src/plugins/hf3fs/hf3fs_backend.cpp b/src/plugins/hf3fs/hf3fs_backend.cpp index de38b83638..dc791710ed 100644 --- a/src/plugins/hf3fs/hf3fs_backend.cpp +++ b/src/plugins/hf3fs/hf3fs_backend.cpp @@ -28,12 +28,15 @@ #include "file/file_utils.h" #define NUM_CQES 1024 +#define HF3FS_DEFAULT_IOPOOL_SIZE 64 +#define HF3FS_MAX_IOPOOL_SIZE (1 << 20) long nixlHf3fsEngine::page_size = sysconf(_SC_PAGESIZE); nixlHf3fsEngine::nixlHf3fsEngine(const nixlBackendInitParams *init_params) : nixlBackendEngine(init_params), - mem_config(NIXL_HF3FS_MEM_CONFIG_AUTO) { + mem_config(NIXL_HF3FS_MEM_CONFIG_AUTO), + iopool_size(HF3FS_DEFAULT_IOPOOL_SIZE) { hf3fs_utils = new hf3fsUtil(); this->initErr = false; @@ -61,6 +64,18 @@ nixlHf3fsEngine::nixlHf3fsEngine(const nixlBackendInitParams *init_params) return; } } + if (init_params->customParams->count("iopool_size") > 0) { + int size = atoi(init_params->customParams->at("iopool_size").c_str()); + if (size > 0) { + if (size < HF3FS_MAX_IOPOOL_SIZE) { + iopool_size = size; + } else { + iopool_size = HF3FS_MAX_IOPOOL_SIZE; + NIXL_WARN << size << " exceeded max iopool size " << iopool_size + << ", set it to max"; + } + } + } } char mount_point_cstr[256]; @@ -72,9 +87,64 @@ nixlHf3fsEngine::nixlHf3fsEngine(const nixlBackendInitParams *init_params) hf3fs_utils->mount_point = mount_point_cstr; - NIXL_DEBUG << "HF3FS: Page size: " << page_size; + for (unsigned int i = 0; i < iopool_size; i++) { + auto io = new nixlHf3fsIO(); + if (io == NULL) { + this->initErr = true; + // io obj will be free when engine destroyed + return; + } + iopool.push_back(io); + } + + NIXL_DEBUG << "HF3FS: page size " << page_size << " iopool_size " << iopool_size; } +nixlHf3fsIO * +nixlHf3fsEngine::getFromIOPool() const { + const std::lock_guard lock(iopool_lock); + if (!iopool.empty()) { + auto io = iopool.front(); + iopool.pop_front(); + return io; + } + return nullptr; +} + +bool +nixlHf3fsEngine::returnToIOPool(nixlHf3fsIO *io) const { + const std::lock_guard lock(iopool_lock); + if (iopool.size() < iopool_size) { + iopool.push_back(io); + return true; + } + return false; +} + +void +nixlHf3fsEngine::destroyIOPool() { + for (auto io : iopool) { + delete (io); + } + iopool.clear(); +} + +nixlHf3fsIO * +nixlHf3fsEngine::getIOObj() const { + auto io_obj = getFromIOPool(); + if (io_obj != nullptr) { + return io_obj; + } + return new nixlHf3fsIO(); +} + +void +nixlHf3fsEngine::putIOObj(nixlHf3fsIO *io) const { + if (returnToIOPool(io)) { + return; + } + delete io; +} nixl_status_t nixlHf3fsEngine::registerMem (const nixlBlobDesc &mem, const nixl_mem_t &nixl_mem, @@ -164,7 +234,7 @@ void nixlHf3fsEngine::cleanupIOList(nixlHf3fsBackendReqH *handle) const if (prev_io->mem_type == NIXL_HF3FS_MEM_TYPE_DRAM) { hf3fs_utils->destroyIOV(&prev_io->iov); } - delete prev_io; + putIOObj(prev_io); } handle->io_list.clear(); @@ -237,10 +307,10 @@ nixl_status_t nixlHf3fsEngine::prepXfer (const nixl_xfer_op_t &operation, offset = (size_t) (*file_list)[i].addr; // Offset in file auto mem_md = (nixlHf3fsMetadata *)(*mem_list)[i].metadataP; - nixlHf3fsIO *io = new nixlHf3fsIO(); + nixlHf3fsIO *io = getIOObj(); if (io == nullptr) { nixl_err = NIXL_ERR_BACKEND; - nixl_mesg = "Error: Failed to create IO"; + nixl_mesg = "Error: Failed to get IO Object"; goto cleanup_handle; } @@ -252,7 +322,7 @@ nixl_status_t nixlHf3fsEngine::prepXfer (const nixl_xfer_op_t &operation, size, shm_md->uuid.get_data().data()); if (status != NIXL_SUCCESS) { - delete io; + putIOObj(io); nixl_err = status; nixl_mesg = "Error: Failed to wrap memory as IOV"; goto cleanup_handle; @@ -260,7 +330,7 @@ nixl_status_t nixlHf3fsEngine::prepXfer (const nixl_xfer_op_t &operation, } else { status = hf3fs_utils->createIOV(&io->iov, size, size); if (status != NIXL_SUCCESS) { - delete io; + putIOObj(io); nixl_err = status; nixl_mesg = "Error: Failed to create IOV"; goto cleanup_handle; @@ -444,6 +514,7 @@ nixl_status_t nixlHf3fsEngine::releaseReqH(nixlBackendReqH* handle) const } nixlHf3fsEngine::~nixlHf3fsEngine() { + destroyIOPool(); hf3fs_utils->closeHf3fsDriver(); delete hf3fs_utils; } @@ -521,7 +592,7 @@ nixlHf3fsDramZCMetadata::nixlHf3fsDramZCMetadata(uint8_t *addr, size_t len, hf3f // Close the file descriptor as it's no longer needed after mmap close(shm_fd); - NIXL_INFO << "Created POSIX shared memory: " << shm_name << " with size: " << len; + NIXL_INFO << "Created shared memory: " << shm_name << " with size: " << len; } nixlHf3fsDramZCMetadata::~nixlHf3fsDramZCMetadata() { @@ -533,5 +604,5 @@ nixlHf3fsDramZCMetadata::~nixlHf3fsDramZCMetadata() { NIXL_PERROR << "Failed to unlink shared memory"; } - NIXL_INFO << "Cleaned up POSIX shared memory: " << shm_name; + NIXL_INFO << "Cleaned up shared memory: " << shm_name; } diff --git a/src/plugins/hf3fs/hf3fs_backend.h b/src/plugins/hf3fs/hf3fs_backend.h index 3f8b0b34cb..ac544904a5 100644 --- a/src/plugins/hf3fs/hf3fs_backend.h +++ b/src/plugins/hf3fs/hf3fs_backend.h @@ -119,6 +119,21 @@ class nixlHf3fsEngine : public nixlBackendEngine { nixl_hf3fs_mem_config mem_config; static long page_size; + mutable std::mutex iopool_lock; + mutable std::list iopool; + unsigned int iopool_size; + + nixlHf3fsIO * + getFromIOPool() const; + bool + returnToIOPool(nixlHf3fsIO *io) const; + void + destroyIOPool(); + nixlHf3fsIO * + getIOObj() const; + void + putIOObj(nixlHf3fsIO *io) const; + void cleanupIOList(nixlHf3fsBackendReqH *handle) const; void cleanupIOThread(nixlHf3fsBackendReqH *handle) const; static void waitForIOsThread(void* handle, void *utils); @@ -138,9 +153,6 @@ class nixlHf3fsEngine : public nixlBackendEngine { bool supportsLocal () const { return true; } - bool supportsProgTh () const { - return false; - } nixl_mem_list_t getSupportedMems () const { nixl_mem_list_t mems; diff --git a/src/plugins/hf3fs/hf3fs_plugin.cpp b/src/plugins/hf3fs/hf3fs_plugin.cpp index 3886b85491..7475aa1985 100644 --- a/src/plugins/hf3fs/hf3fs_plugin.cpp +++ b/src/plugins/hf3fs/hf3fs_plugin.cpp @@ -18,70 +18,24 @@ #include "backend/backend_plugin.h" #include "hf3fs_backend.h" #include -// Plugin version information -static const char* PLUGIN_NAME = "HF3FS"; -static const char* PLUGIN_VERSION = "0.1.0"; -// Function to create a new HF3FS backend engine instance -static nixlBackendEngine* create_hf3fs_engine(const nixlBackendInitParams* init_params) { - return new nixlHf3fsEngine(init_params); -} - -static void destroy_hf3fs_engine(nixlBackendEngine *engine) { - delete engine; -} - -// Function to get the plugin name -static const char* get_plugin_name() { - return PLUGIN_NAME; -} - -// Function to get the plugin version -static const char* get_plugin_version() { - return PLUGIN_VERSION; -} - -// Function to get backend options -static nixl_b_params_t get_backend_options() { - nixl_b_params_t params; - return params; -} - -// Function to get supported backend mem types -static nixl_mem_list_t get_backend_mems() { - nixl_mem_list_t mems; - mems.push_back(FILE_SEG); - mems.push_back(DRAM_SEG); - return mems; -} -// Static plugin structure -static nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_hf3fs_engine, - destroy_hf3fs_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems -}; +// Plugin type alias for convenience +using hf3fs_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_HF3FS - -nixlBackendPlugin* createStaticHf3fsPlugin() { - return &plugin; // Return the static plugin instance +nixlBackendPlugin * +createStaticHF3FSPlugin() { + return hf3fs_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "HF3FS", "0.1.0", {}, {FILE_SEG, DRAM_SEG}); } - #else - -// Plugin initialization function -extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - return &plugin; -} - -// Plugin cleanup function -extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { - // Cleanup any resources if needed +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return hf3fs_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "HF3FS", "0.1.0", {}, {FILE_SEG, DRAM_SEG}); } +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} #endif diff --git a/src/plugins/hf3fs/meson.build b/src/plugins/hf3fs/meson.build index a1fbb71560..f32993cd9b 100644 --- a/src/plugins/hf3fs/meson.build +++ b/src/plugins/hf3fs/meson.build @@ -40,7 +40,8 @@ else install: true, cpp_args : ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', 'echo "HF3FS=' + hf3fs_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', diff --git a/src/plugins/libfabric/README.md b/src/plugins/libfabric/README.md new file mode 100644 index 0000000000..2846bfedb2 --- /dev/null +++ b/src/plugins/libfabric/README.md @@ -0,0 +1,85 @@ +#NIXL Libfabric Plugin + +This plugin provides a high-performance RDMA backend for NIXL using the OpenFabrics Interfaces (OFI) Libfabric library. + +## Overview + +The Libfabric plugin provides a high-performance RDMA communication backend with the following key capabilities: + +- **Multi-Rail RDMA**: Automatic discovery and utilization of multiple network devices for increased bandwidth +- **GPU Direct Support**: Zero-copy transfers between GPU memory (VRAM) and remote systems with CUDA integration. And GDR support is currently mandated +- **Scalable Connection Management**: Efficient multi-agent connectivity with robust state tracking and automatic reconnection +- **Asynchronous Processing**: Non-blocking RDMA operations with pre-allocated request pools and completion processing +- **Thread-Safe Concurrency**: Background progress threads with lock-free data structures and configurable threading patterns + +EFA Specific **Topology-Aware Optimization**: Hardware-aware GPU-to-EFA and NUMA-to-EFA mapping using hwloc for optimal performance + +## Dependencies + +### Required Dependencies + +- **Libfabric** + - Many system will have installed libfabric already. If not, custom libfabric installation is available via https://ofiwg.github.io/libfabric/ - Minimum required version: v2.3.0rc2 + - For EFA enabled AWS instances, it is recommanded to install through AWS EFA installer: https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/efa-start.html - Minimum required version: 1.43.2 + +- **hwloc** + - hwloc is used to understand the underlying architecture to optimize application performance. Suggested version: 2.10.0 or newer + +### Network Hardware Requirements + +Validated compatiblity with: +- **AWS EFA** (Elastic Fabric Adapter) + +Any other Libfabric providers that support heterogeneous memory (FI_HMEM) should also work but have not been validated in production environments. Community validation and feedback are highly appreciated! + +## Build Instructions + +```bash +#Basic build setup with default options +$ meson setup + +#Setup with custom options(example) +$ meson setup \ + -Dlibfabric_path=/path/to/libfabric + +#Build and install +ninja && ninja install +``` + +## API Reference + +### Core Classes + +- **`nixlLibfabricEngine`** - Main backend engine providing multi-rail RDMA operations with GPU Direct support +- **`nixlLibfabricRailManager`** - Manages multiple network rails with topology-aware selection and striping strategies +- **`nixlLibfabricRail`** - Individual network rail handling libfabric resources and completion processing +- **`nixlLibfabricTopology`** - Hardware topology discovery for optimal GPU-to-EFA and NUMA-to-EFA mapping +- **`nixlLibfabricBackendH`** - Request handle for tracking multi-request transfer completion with atomic counters +- **`nixlLibfabricConnection`** - Multi-rail connection metadata for remote agents with state management + +## Troubleshooting + +### Debug Information + +Enable debug logging by setting environment variables: +```bash +#Libfabric debug logging +export FI_LOG_LEVEL=debug +export FI_LOG_PROV=efa # or verbs, tcp, etc. + +#NIXL debug logging +export NIXL_LOG_LEVEL=debug +``` + +### Common Issues + +**No network devices detected:** +```bash +#Check available fabric interfaces +fi_info -l + +#For checking specific devices(e.g.EFA as an example) +fi_info -p efa +``` + +For additional support, check the NIXL documentation and Libfabric provider-specific guides. diff --git a/src/plugins/libfabric/libfabric_backend.cpp b/src/plugins/libfabric/libfabric_backend.cpp new file mode 100644 index 0000000000..c5216ac1e4 --- /dev/null +++ b/src/plugins/libfabric/libfabric_backend.cpp @@ -0,0 +1,1557 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric_backend.h" +#include "serdes/serdes.h" +#include "common/nixl_log.h" +#include "libfabric/libfabric_topology.h" + +#include +#include +#include + +#include +#include + +#include "absl/strings/numbers.h" + +#ifdef HAVE_CUDA +// CUDA error checking macros +#define CHECK_CUDA_ERROR(result, message) \ + do { \ + if (result != cudaSuccess) { \ + NIXL_ERROR << "CUDA Error: " << message << " (" << cudaGetErrorString(result) << ")"; \ + return NIXL_ERR_BACKEND; \ + } \ + } while (0) + +#define CHECK_CUDA_DRIVER_ERROR(result, message) \ + do { \ + if (result != CUDA_SUCCESS) { \ + const char *error_str; \ + cuGetErrorString(result, &error_str); \ + NIXL_ERROR << "CUDA Driver Error: " << message << " (" << error_str << ")"; \ + return NIXL_ERR_BACKEND; \ + } \ + } while (0) +#endif + +/**************************************** + * CUDA Context Management + *****************************************/ + +#ifdef HAVE_CUDA +static int +cudaQueryAddr(void *address, bool &is_dev, CUdevice &dev, CUcontext &ctx) { + CUmemorytype mem_type = CU_MEMORYTYPE_HOST; + uint32_t is_managed = 0; + CUpointer_attribute attr_type[4]; + void *attr_data[4]; + CUresult result; + + attr_type[0] = CU_POINTER_ATTRIBUTE_MEMORY_TYPE; + attr_data[0] = &mem_type; + attr_type[1] = CU_POINTER_ATTRIBUTE_IS_MANAGED; + attr_data[1] = &is_managed; + attr_type[2] = CU_POINTER_ATTRIBUTE_DEVICE_ORDINAL; + attr_data[2] = &dev; + attr_type[3] = CU_POINTER_ATTRIBUTE_CONTEXT; + attr_data[3] = &ctx; + + result = cuPointerGetAttributes(4, attr_type, attr_data, (CUdeviceptr)address); + is_dev = (mem_type == CU_MEMORYTYPE_DEVICE); + + return (CUDA_SUCCESS != result); +} + +void +nixlLibfabricCudaCtx::cudaResetCtxPtr() { + pthrCudaCtx_ = NULL; + myDevId_ = -1; +} + +int +nixlLibfabricCudaCtx::cudaUpdateCtxPtr(void *address, int expected_dev, bool &was_updated) { + bool is_dev; + CUdevice dev; + CUcontext ctx; + int ret; + + was_updated = false; + + if (expected_dev == -1) return -1; + if (myDevId_ != -1 && expected_dev != myDevId_) return -1; + + ret = cudaQueryAddr(address, is_dev, dev, ctx); + if (ret) return ret; + if (!is_dev) return 0; + if (dev != expected_dev) return -1; + + if (pthrCudaCtx_) { + if (pthrCudaCtx_ != ctx) return -1; + return 0; + } + + pthrCudaCtx_ = ctx; + was_updated = true; + myDevId_ = expected_dev; + + return 0; +} + +int +nixlLibfabricCudaCtx::cudaSetCtx() { + CUresult result; + if (NULL == pthrCudaCtx_) return 0; + + result = cuCtxSetCurrent(pthrCudaCtx_); + return (CUDA_SUCCESS == result); +} + +void +nixlLibfabricEngine::vramInitCtx() { + cudaCtx_ = std::make_unique(); +} + +int +nixlLibfabricEngine::vramUpdateCtx(void *address, uint64_t devId, bool &restart_reqd) { + int ret; + bool was_updated; + + restart_reqd = false; + + if (!cuda_addr_wa_) { + return 0; // Nothing to do + } + + ret = cudaCtx_->cudaUpdateCtxPtr(address, devId, was_updated); + if (ret) { + return ret; + } + + restart_reqd = was_updated; + return 0; +} + +int +nixlLibfabricEngine::vramApplyCtx() { + if (!cuda_addr_wa_) { + return 0; // Nothing to do + } + return cudaCtx_->cudaSetCtx(); +} + +void +nixlLibfabricEngine::vramFiniCtx() { + cudaCtx_.reset(); +} +#endif + +/**************************************** + * Request Management + *****************************************/ + +nixlLibfabricBackendH::nixlLibfabricBackendH() : completed_requests_(0), total_requests_used_(0) { + NIXL_DEBUG << "constructor called, this: " << this + << " total_requests_used=" << total_requests_used_.load(); +} + +nixlLibfabricBackendH::~nixlLibfabricBackendH() { + NIXL_DEBUG << "destructor called, this: " << this; +} + +// Multi-request completion tracking methods +void +nixlLibfabricBackendH::init_request_tracking(size_t num_requests) { + total_requests_used_.store(num_requests); + completed_requests_.store(0); + NIXL_DEBUG << "Initialized request tracking for " << num_requests << " requests"; +} + +void +nixlLibfabricBackendH::increment_completed_requests() { + size_t completed = completed_requests_.fetch_add(1); + NIXL_DEBUG << "Request completed, total completed: " << completed << "/" + << total_requests_used_.load(); +} + +size_t +nixlLibfabricBackendH::get_completed_requests_count() const { + return completed_requests_.load(); +} + +size_t +nixlLibfabricBackendH::get_total_requests_used() const { + return total_requests_used_.load(); +} + +void +nixlLibfabricBackendH::adjust_total_requests(size_t actual_count) { + total_requests_used_.store(actual_count); + NIXL_DEBUG << "Adjusted total requests to actual count: " << actual_count; +} + +bool +nixlLibfabricBackendH::is_completed() const { + // Transfer is completed when all requests have completed + // NIXL_DEBUG << "Request completed, total completed: " << completed_requests_.load(); + return completed_requests_.load() == total_requests_used_.load(); +} + +/**************************************** + * Constructor/Destructor + *****************************************/ + +nixlLibfabricEngine::nixlLibfabricEngine(const nixlBackendInitParams *init_params) + : nixlBackendEngine(init_params), + cm_thread_stop_(false), + progress_thread_enabled_(init_params->enableProgTh), + progress_thread_delay_(std::chrono::microseconds(init_params->pthrDelay)), + rail_manager(NIXL_LIBFABRIC_DEFAULT_STRIPING_THRESHOLD) { + + NIXL_DEBUG << "Initializing Libfabric Backend with GPU Support"; + +#ifdef HAVE_CUDA + // Initialize CUDA context management + vramInitCtx(); + // CUDA address workaround + if (getenv("NIXL_DISABLE_CUDA_ADDR_WA")) { + NIXL_DEBUG << "Disabling CUDA address workaround"; + cuda_addr_wa_ = false; + } else { + cuda_addr_wa_ = true; + NIXL_DEBUG << "CUDA address workaround enabled"; + } +#endif + + // Parse striping threshold parameter + std::string threshold_str; + striping_threshold_ = NIXL_LIBFABRIC_DEFAULT_STRIPING_THRESHOLD; + + if (getInitParam("striping_threshold", threshold_str) == NIXL_SUCCESS) { + try { + striping_threshold_ = std::stoull(threshold_str); + NIXL_DEBUG << "Using custom striping threshold: " << striping_threshold_ << " bytes"; + } + catch (const std::exception &e) { + NIXL_WARN << "Invalid striping_threshold value '" << threshold_str + << "', using default: " << striping_threshold_ << " bytes"; + } + } else { + NIXL_DEBUG << "Using default striping threshold: " << striping_threshold_ << " bytes"; + } + + // Initialize Rail Manager which will discover the topology and create all rails. + try { + NIXL_DEBUG << "Rail Manager created with " << rail_manager.getNumDataRails() + << " data rails and " << rail_manager.getNumControlRails() << " control rails"; + + // Set up callbacks on each rail using Engine's static callback functions + size_t control_rail_id = 0; + NIXL_DEBUG << "Set notification processor for control rail 0"; + rail_manager.getControlRail(control_rail_id) + .setNotificationCallback([this](const std::string &serialized_notif) { + processNotification(serialized_notif); + }); + + // Set up connection state callbacks for control rails + NIXL_DEBUG << "Set connection state processor for CM rail 0"; + + rail_manager.getControlRail(control_rail_id) + .setConnectionAckCallback([this](const uint16_t agent_idx, + nixlLibfabricConnection *conn_info, + ConnectionState state) { + processConnectionAck(agent_idx, conn_info, state); + }); + + // Set up connection request callback for control rails + rail_manager.getControlRail(control_rail_id) + .setConnectionReqCallback([this](const uint16_t agent_idx, + const std::string &serialized_data, + nixlLibfabricRail *rail) -> nixl_status_t { + return processConnectionRequest(agent_idx, serialized_data, rail); + }); + + // Set up XFER_ID tracking callbacks for all data rails + NIXL_DEBUG << "Setting up XFER_ID tracking callbacks for " << rail_manager.getNumDataRails() + << " data rails"; + for (size_t data_rail_id = 0; data_rail_id < rail_manager.getNumDataRails(); + ++data_rail_id) { + rail_manager.getDataRail(data_rail_id).setXferIdCallback([this](uint32_t xfer_id) { + addReceivedXferId(xfer_id); + }); + NIXL_DEBUG << "Set XFER_ID callback for data rail " << data_rail_id; + } + + // Create self-connection + std::vector> data_endpoints( + rail_manager.getNumDataRails()); + std::vector> control_endpoints( + rail_manager.getNumControlRails()); + // Prepare data rail endpoints + for (size_t rail_id = 0; rail_id < rail_manager.getNumDataRails(); ++rail_id) { + std::memcpy(data_endpoints[rail_id].data(), + rail_manager.getDataRail(rail_id).ep_name, + sizeof(rail_manager.getDataRail(rail_id).ep_name)); + } + // Prepare control rail endpoints + for (size_t rail_id = 0; rail_id < rail_manager.getNumControlRails(); ++rail_id) { + std::memcpy(control_endpoints[rail_id].data(), + rail_manager.getControlRail(rail_id).ep_name, + sizeof(rail_manager.getControlRail(rail_id).ep_name)); + } + // Create self-connection using common method + nixl_status_t conn_status = + createAgentConnection(localAgent, data_endpoints, control_endpoints); + if (conn_status != NIXL_SUCCESS) { + throw std::runtime_error( + "createAgentConnection failed for self-connection with status: " + + std::to_string(conn_status)); + } + + NIXL_DEBUG << "Created self-connection for agent: " << localAgent << " on " + << rail_manager.getNumDataRails() << " data rails and " + << rail_manager.getNumControlRails() << " control rails"; + + // Threading infrastructure + // Start CM thread for background processing + NIXL_DEBUG << "Starting CM thread"; + cm_thread_ = std::thread(&nixlLibfabricEngine::cmThread, this); + if (!cm_thread_.joinable()) { + NIXL_ERROR << "Failed to start CM thread"; + throw std::runtime_error("Failed to start CM thread"); + } + NIXL_DEBUG << "ConnectionManagement thread started successfully"; + + // Start Progress thread for data rail completion processing + if (progress_thread_enabled_) { + NIXL_DEBUG << "Starting Progress thread for data rails with delay: " + << progress_thread_delay_.count() << " microseconds"; + progress_thread_stop_ = false; + progress_thread_ = std::thread(&nixlLibfabricEngine::progressThread, this); + + if (!progress_thread_.joinable()) { + NIXL_ERROR << "Failed to start Progress thread"; + throw std::runtime_error("Failed to start Progress thread"); + } + NIXL_DEBUG << "Progress thread started successfully"; + } else { + NIXL_DEBUG << "Progress thread disabled, using manual progress in checkXfer/getNotifs"; + } + } + catch (const std::exception &e) { + cleanup(); + throw; + } +} + +nixlLibfabricEngine::~nixlLibfabricEngine() { + NIXL_DEBUG + << "Destructor starting, stopping all threads FIRST to prevent timing report interruption"; + + // STOP ALL THREADS FIRST to prevent any interference with timing report + cm_thread_stop_.store(true); + + if (progress_thread_enabled_) { + progress_thread_stop_.store(true); + } + + // Post dummy completion to wake up blocking threads + postShutdownCompletion(); + + if (cm_thread_.joinable()) { + NIXL_DEBUG << "Waiting for CM thread to exit"; + cm_thread_.join(); + NIXL_DEBUG << "CM thread joined successfully"; + } + if (progress_thread_enabled_ && progress_thread_.joinable()) { + NIXL_DEBUG << "Waiting for Progress thread to exit"; + progress_thread_.join(); + NIXL_DEBUG << "Progress thread joined successfully"; + } else if (!progress_thread_enabled_) { + NIXL_DEBUG << "Progress thread was not running"; + } + NIXL_DEBUG << "All threads stopped, now cleaning up resources"; + cleanup(); +} + +/**************************************** + * Connection management + *****************************************/ + +nixl_status_t +nixlLibfabricEngine::getConnInfo(std::string &str) const { + // Verify all rail endpoints are initialized + for (size_t rail_id = 0; rail_id < rail_manager.getNumDataRails(); ++rail_id) { + if (!rail_manager.getDataRail(rail_id).endpoint) { + NIXL_ERROR << "Rail " << rail_id << " endpoint not initialized"; + return NIXL_ERR_BACKEND; + } + } + + NIXL_DEBUG << "Retrieving local endpoint addresses for all " << rail_manager.getNumDataRails() + << " rails"; + + // Use Rail Manager's connection SerDes method with "dest" prefix for remote consumption + nixl_status_t status = rail_manager.serializeConnectionInfo("dest", str); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail Manager serializeConnectionInfo failed"; + return status; + } + + NIXL_DEBUG << "Rail Manager serialized connection info for " << rail_manager.getNumDataRails() + << " rails, " << rail_manager.getNumControlRails() << " control rails, " + << "total size: " << str.length(); + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::loadRemoteConnInfo(const std::string &remote_agent, + const std::string &remote_conn_info) { + std::lock_guard lock(connection_state_mutex_); + + NIXL_DEBUG << "Loading remote info for agent: " << remote_agent + << ", info length: " << remote_conn_info.length() + << ", info (hex): " << LibfabricUtils::hexdump(remote_conn_info.data()); + + if (remote_conn_info.empty()) { + NIXL_ERROR << "Empty remote connection info received"; + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_DEBUG << "Processing " << rail_manager.getNumDataRails() << " data rails and " + << rail_manager.getNumControlRails() << " control rails for agent: " << remote_agent; + + // Use Rail Manager's connection SerDes method with "dest" prefix (remote is sending us their + // endpoints as "dest") + std::vector> data_endpoints; + std::vector> control_endpoints; + nixl_status_t status = rail_manager.deserializeConnectionInfo( + "dest", remote_conn_info, data_endpoints, control_endpoints); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail Manager deserializeConnectionInfo failed"; + return status; + } + // Create connection to remote agent + nixl_status_t conn_status = + createAgentConnection(remote_agent, data_endpoints, control_endpoints); + if (conn_status != NIXL_SUCCESS) { + NIXL_ERROR << "createAgentConnection failed with status: " << conn_status; + return conn_status; + } + + NIXL_DEBUG << "Successfully stored multirail connection for " << remote_agent << " on " + << rail_manager.getNumDataRails() << " rails"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::connect(const std::string &remote_agent) { + std::lock_guard lock(connection_state_mutex_); + + NIXL_DEBUG << "Connecting to agent: " << remote_agent + << ", connections_ size: " << connections_.size(); + + // Check if connection is already established + auto it = connections_.find(remote_agent); + if (it != connections_.end() && it->second->overall_state_ == ConnectionState::CONNECTED) { + NIXL_DEBUG << "Connection already established for " << remote_agent + << ", fi_addr: " << it->second->rail_remote_addr_list_[0]; + return NIXL_SUCCESS; + } + + // Connection exists but not established - trigger establishConnection() + NIXL_DEBUG << "Connection exists but not established, triggering establishConnection for " + << remote_agent; + + // Release the lock before calling establishConnection since it acquires the same mutex + lock.~lock_guard(); + + nixl_status_t status = establishConnection(remote_agent); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to establish connection with " << remote_agent; + return status; + } + + it = connections_.find(remote_agent); + if (it == connections_.end()) { + NIXL_DEBUG << "Connect failed. No metadata connection info for " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + NIXL_DEBUG << "Successfully established connection for " << remote_agent; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::disconnect(const std::string &remote_agent) { + std::lock_guard lock(connection_state_mutex_); + auto it = connections_.find(remote_agent); + if (it == connections_.end()) { + NIXL_ERROR << "Disconnect failed. No metadata connection info for " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + // Connection exists - check if already disconnected + if (it->second->overall_state_ == ConnectionState::DISCONNECTED) { + NIXL_DEBUG << "Connection already established for " << remote_agent + << ", fi_addr: " << it->second->rail_remote_addr_list_[0]; + return NIXL_SUCCESS; + } + // TODO: Implement disconnect logic to cleanup the AV Address Entries from both local and remote + // AV. + + // Update connection state to DISCONNECTED before removing + it->second->overall_state_ = ConnectionState::DISCONNECTED; + + // Remove connection from map + connections_.erase(remote_agent); + NIXL_DEBUG << "Connection erased from the connection map for agent: " << remote_agent; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::createAgentConnection( + const std::string &agent_name, + const std::vector> &data_rail_endpoints, + const std::vector> &control_rail_endpoints) { + + NIXL_DEBUG << "Creating connection for agent: " << agent_name; + + // Validate input parameters + if (data_rail_endpoints.size() != rail_manager.getNumDataRails()) { + NIXL_ERROR << "Expected " << rail_manager.getNumDataRails() << " data rail endpoints, got " + << data_rail_endpoints.size(); + return NIXL_ERR_INVALID_PARAM; + } + + if (control_rail_endpoints.size() != rail_manager.getNumControlRails()) { + NIXL_ERROR << "Expected " << rail_manager.getNumControlRails() + << " control rail endpoints, got " << control_rail_endpoints.size(); + return NIXL_ERR_INVALID_PARAM; + } + + // Create connection object + auto conn = std::make_shared(); + if (!conn) { + NIXL_ERROR << "Failed to allocate connection object"; + return NIXL_ERR_BACKEND; + } + + conn->remoteAgent_ = agent_name; + conn->rail_remote_addr_list_.reserve(rail_manager.getNumDataRails()); + conn->control_rail_remote_addr_list_.reserve(rail_manager.getNumControlRails()); + + // Process all data rails in one operation + nixl_status_t data_status = + rail_manager.insertAllAddresses(nixlLibfabricRailManager::RailType::DATA, + data_rail_endpoints, + conn->rail_remote_addr_list_, + conn->src_ep_names_); + if (data_status != NIXL_SUCCESS) { + NIXL_ERROR << "insertAllAddresses failed for data rails with status: " << data_status; + return NIXL_ERR_BACKEND; + } + + // Process all control rails in one operation + nixl_status_t control_status = + rail_manager.insertAllAddresses(nixlLibfabricRailManager::RailType::CONTROL, + control_rail_endpoints, + conn->control_rail_remote_addr_list_, + conn->control_ep_names_); + if (control_status != NIXL_SUCCESS) { + NIXL_ERROR << "insertAllAddresses failed for control rails with status: " << control_status; + return NIXL_ERR_BACKEND; + } + + // Manage agent names and index + agent_names_.push_back(agent_name); + int index = 0; + std::for_each(agent_names_.begin(), agent_names_.end(), [&index](const std::string &name) { + NIXL_DEBUG << "Index " << index << ": " << name; + index++; + }); + conn->agent_index_ = agent_names_.size() - 1; + + // Store connection + connections_[agent_name] = conn; + + NIXL_DEBUG << "Successfully created connection for agent: " << agent_name << " on " + << rail_manager.getNumDataRails() << " data rails and " + << rail_manager.getNumControlRails() << " control rails"; + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::establishConnection(const std::string &remote_agent) const { + // Use existing connection_state_mutex_ to serialize connection establishment + std::lock_guard lock(connection_state_mutex_); + + // Check if another thread already established the connection + auto it = connections_.find(remote_agent); + if (it != connections_.end() && it->second->overall_state_ == ConnectionState::CONNECTED) { + NIXL_DEBUG << "Connection already established by another thread for " << remote_agent; + return NIXL_SUCCESS; + } + + if (it == connections_.end()) { + NIXL_ERROR << "No connection found for agent: " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + + // Verify we have addresses for all data rails + if (it->second->rail_remote_addr_list_.size() != rail_manager.getNumDataRails()) { + NIXL_ERROR << "Remote connection has " << it->second->rail_remote_addr_list_.size() + << " data rails, expected " << rail_manager.getNumDataRails(); + return NIXL_ERR_BACKEND; + } + + NIXL_DEBUG << "Establishing connections_ on control rails and data rails for agent: " + << remote_agent; + + // Use single "Communicator" for CM + auto *conn_info = reinterpret_cast(it->second.get()); + + NIXL_DEBUG << "Using connection info : 0: " + << LibfabricUtils::hexdump(conn_info->src_ep_names_[0]) << std::endl + << "1: " << LibfabricUtils::hexdump(conn_info->src_ep_names_[1]) << std::endl + << "control_0: " << LibfabricUtils::hexdump(conn_info->control_ep_names_[0]) + << std::endl + << " with agent index: " << it->second->agent_index_; + if (!conn_info) { + NIXL_ERROR << "Connection info for agent " << remote_agent << " is null"; + return NIXL_ERR_BACKEND; + } + + // Allocate control request + const size_t control_rail_id = 0; + + // Serialize connection info + std::string serialized_conn_info; + nixl_status_t serialize_status = + rail_manager.serializeConnectionInfo("src", serialized_conn_info); + if (serialize_status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail manager serializeConnectionInfo failed"; + return serialize_status; + } + + nixlLibfabricReq *control_request = rail_manager.getControlRail(control_rail_id) + .allocateControlRequest(serialized_conn_info.length()); + if (!control_request) { + NIXL_ERROR << "Failed to allocate control request for connection establishment"; + return NIXL_ERR_BACKEND; + } + + // Copy serialized data to control request buffer + memcpy(control_request->buffer, serialized_conn_info.data(), serialized_conn_info.length()); + control_request->buffer_size = serialized_conn_info.length(); + + nixl_status_t status = rail_manager.postControlMessage( + nixlLibfabricRailManager::ControlMessageType::CONNECTION_REQ, + control_request, + conn_info->control_rail_remote_addr_list_[0], // Always use control rail 0 + it->second->agent_index_ // agent_index is only used in the ACK back from remote, + // to match connection request + ); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "postSend failed on rail " << 0; + // TODO, wrap req info into a nixlLibfabricRequestHandle and add retry logic + return NIXL_ERR_BACKEND; + } + // Register the connection state tracker with the CM thread + // Wait for the CM thread to establish the connection + // TODO: Currently blocking, update to timeout and return NIXL_IN_PROG + { + std::unique_lock lock(conn_info->conn_state_mutex_); + NIXL_DEBUG << "Waiting for connection to be established for agent: " << remote_agent; + conn_info->cv_.wait(lock, [conn_info] { + return conn_info->overall_state_ == ConnectionState::CONNECTED || + conn_info->overall_state_ == ConnectionState::FAILED; + }); + NIXL_DEBUG << "Connection state for agent " << remote_agent << " is now " + << conn_info->overall_state_; + + if (conn_info->overall_state_ == ConnectionState::FAILED) { + NIXL_ERROR << "Connection failed on control rail 0"; + return NIXL_ERR_BACKEND; + } + } + + NIXL_DEBUG << "Connection already established for agent: " << remote_agent; + return NIXL_SUCCESS; +} + +/**************************************** + * Memory management + *****************************************/ + +nixl_mem_list_t +nixlLibfabricEngine::getSupportedMems() const { + nixl_mem_list_t mems; + mems.push_back(DRAM_SEG); +#ifdef HAVE_CUDA + mems.push_back(VRAM_SEG); +#endif + return mems; +} + +nixl_status_t +nixlLibfabricEngine::registerMem(const nixlBlobDesc &mem, + const nixl_mem_t &nixl_mem, + nixlBackendMD *&out) { + auto priv = std::make_unique(); + + priv->buffer_ = (void *)mem.addr; + priv->length_ = mem.len; + priv->gpu_device_id_ = mem.devId; // Store GPU device ID + +#ifdef HAVE_CUDA + // Handle CUDA memory registration with GPU Direct RDMA support + if (nixl_mem == VRAM_SEG) { + // For multi-GPU support, skip CUDA address workaround + if (cuda_addr_wa_) { + bool need_restart; + if (vramUpdateCtx((void *)mem.addr, mem.devId, need_restart)) { + NIXL_WARN << "CUDA address workaround failed for device " << mem.devId + << ", disabling workaround for multi-GPU support"; + cuda_addr_wa_ = false; // Disable workaround for subsequent registrations + } else if (need_restart) { + // Restart progress thread if needed + NIXL_DEBUG << "CUDA context updated, restarting progress thread"; + vramApplyCtx(); + } + } + // Set CUDA device context directly for multi-GPU support + if (!cuda_addr_wa_) { + cudaError_t cuda_ret = cudaSetDevice(mem.devId); + if (cuda_ret != cudaSuccess) { + NIXL_ERROR << "Failed to set CUDA device " << mem.devId << ": " + << cudaGetErrorString(cuda_ret); + return NIXL_ERR_NOT_SUPPORTED; + } + NIXL_DEBUG << "Set CUDA device context to GPU " << mem.devId; + } + } +#endif + + // Initialize vectors to accommodate all possible rails (for indexing consistency) + priv->rail_mr_list_.resize(rail_manager.getNumDataRails(), nullptr); + priv->rail_key_list_.resize(rail_manager.getNumDataRails(), 0); + +#ifdef HAVE_CUDA + // Set CUDA context before libfabric operations for VRAM + if (nixl_mem == VRAM_SEG) { + vramApplyCtx(); + } +#endif + + // Use Rail Manager for centralized memory registration with GPU Direct RDMA support + nixl_status_t status = rail_manager.registerMemory((void *)mem.addr, + mem.len, + nixl_mem, + mem.devId, + priv->rail_mr_list_, + priv->rail_key_list_, + priv->selected_rails_); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail Manager registerMemory failed"; + return status; + } + + NIXL_DEBUG << "Rail Manager successfully registered " + << (nixl_mem == VRAM_SEG ? "VRAM" : "DRAM") << " memory on " + << priv->selected_rails_.size() << " rails" + << (nixl_mem == VRAM_SEG ? " with GPU Direct RDMA support" : ""); + + NIXL_DEBUG << "Successfully registered memory on " << priv->selected_rails_.size() + << " rails for " << (nixl_mem == VRAM_SEG ? "GPU" : "CPU") << " " << mem.devId; + out = priv.release(); + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::deregisterMem(nixlBackendMD *meta) { + auto *priv = static_cast(meta); + // Use Rail Manager for centralized memory deregistration + nixl_status_t status = + rail_manager.deregisterMemory(priv->selected_rails_, priv->rail_mr_list_); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail Manager deregisterMemory failed"; + // Continue with cleanup even if deregistration failed + } + + delete priv; + return status; +} + +nixl_status_t +nixlLibfabricEngine::getPublicData(const nixlBackendMD *meta, std::string &str) const { + const nixlLibfabricPrivateMetadata *priv = + static_cast(meta); + + return rail_manager.serializeMemoryKeys(priv->rail_key_list_, priv->buffer_, str); +} + +nixl_status_t +nixlLibfabricEngine::loadLocalMD(nixlBackendMD *input, nixlBackendMD *&output) { + nixlLibfabricPrivateMetadata *input_md = static_cast(input); + auto pub_md = std::make_unique(); + // Store all rail keys instead of just the first one + pub_md->rail_remote_key_list_.reserve(input_md->rail_key_list_.size()); + for (size_t rail_id = 0; rail_id < input_md->rail_key_list_.size(); ++rail_id) { + pub_md->rail_remote_key_list_.push_back(input_md->rail_key_list_[rail_id]); + NIXL_DEBUG << "Added rail " << rail_id << " key: " << input_md->rail_key_list_[rail_id]; + } + + pub_md->remote_buf_addr_ = reinterpret_cast(input_md->buffer_); + pub_md->conn_ = connections_[localAgent]; + + output = pub_md.release(); + NIXL_DEBUG << "Loading Local MD with " << input_md->rail_key_list_.size() << " rail keys"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::loadRemoteMD(const nixlBlobDesc &input, + const nixl_mem_t &nixl_mem, + const std::string &remote_agent, + nixlBackendMD *&output) { + NIXL_DEBUG << "Loading remote metadata for agent: " << remote_agent; + + auto conn_it = connections_.find(remote_agent); + if (conn_it == connections_.end()) { + NIXL_ERROR << "Could not find connection for agent: " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + // Delegate to Rail Manager for SerDes operations (returns raw data) + std::vector remote_keys; + uint64_t remote_addr; + nixl_status_t status = + rail_manager.deserializeMemoryKeys(input.metaInfo, remote_keys, remote_addr); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Rail Manager deserializeMemoryKeys failed"; + return status; + } + + // Engine handles connection management and metadata object creation + auto pub_md = std::make_unique(); + pub_md->conn_ = conn_it->second; + pub_md->rail_remote_key_list_ = std::move(remote_keys); + pub_md->remote_buf_addr_ = remote_addr; + NIXL_DEBUG << "Remote metadata loaded with" + << " Remote addr: " << (void *)pub_md->remote_buf_addr_ << " Remote keys for " + << pub_md->rail_remote_key_list_.size() << " rails" + << " Remote fi_addr: " << pub_md->conn_->rail_remote_addr_list_[0]; + + output = pub_md.release(); + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::unloadMD(nixlBackendMD *input) { + delete input; + return NIXL_SUCCESS; +} + +/**************************************** + * Data movement + *****************************************/ + +nixl_status_t +nixlLibfabricEngine::prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { + NIXL_DEBUG << "Preparing transfer for remote_agent: " << remote_agent; + + auto conn_it = connections_.find(remote_agent); + if (conn_it == connections_.end() || !conn_it->second) { + NIXL_ERROR << "No valid connection found for agent: " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + auto backend_handle = new nixlLibfabricBackendH(); + if (!backend_handle) { + NIXL_ERROR << "Failed to allocate nixlLibfabricBackendH"; + return NIXL_ERR_BACKEND; + } + handle = backend_handle; // Assign to base class pointer + + NIXL_DEBUG << "Transfer preparation complete, handle address: " << handle; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::estimateXferCost(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *const &handle, + std::chrono::microseconds &duration, + std::chrono::microseconds &err_margin, + nixl_cost_t &method, + const nixl_opt_args_t *opt_args) const { + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { + + // Validate connection + auto conn_it = connections_.find(remote_agent); + if (conn_it == connections_.end() || !conn_it->second) { + NIXL_ERROR << "No valid connection found for agent: " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + if (conn_it->second->overall_state_ == ConnectionState::DISCONNECTED) { + NIXL_DEBUG << "No existing connection for " << remote_agent + << ", establishing new connection"; + nixl_status_t status = this->establishConnection(remote_agent); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to establish connection with " << remote_agent; + return status; + } + NIXL_DEBUG << "Established new connection with remote_agent: " << remote_agent; + } + + NIXL_DEBUG << "Posting transfer for remote_agent: " << remote_agent + << ", handle address: " << handle; + + auto backend_handle = static_cast(handle); + if (!backend_handle) { + NIXL_ERROR << "Failed to cast handle to nixlLibfabricBackendH"; + return NIXL_ERR_INVALID_PARAM; + } + + // Allocate a new notification request at the start of each postXfer + const size_t control_rail_id = 0; + nixlLibfabricReq *control_request = rail_manager.getControlRail(control_rail_id) + .allocateControlRequest(sizeof(BinaryNotification)); + if (!control_request) { + NIXL_ERROR << "Failed to allocate control request for notification"; + return NIXL_ERR_BACKEND; + } + + // Create BinaryNotification directly in the control request buffer + BinaryNotification *binary_notif = + reinterpret_cast(control_request->buffer); + binary_notif->clear(); + + nixlLibfabricReq::OpType op_type; + int desc_count = local.descCount(); + + NIXL_DEBUG << "Processing " << desc_count + << " descriptors using optimized single-pass approach"; + + op_type = (operation == NIXL_WRITE) ? nixlLibfabricReq::WRITE : nixlLibfabricReq::READ; + + // Set initial request count to maximum possible requests + size_t max_possible_requests = desc_count * rail_manager.getNumDataRails(); + backend_handle->init_request_tracking(max_possible_requests); + + // Core transfer submission to process each descriptor with direct submission + for (int desc_idx = 0; desc_idx < desc_count; ++desc_idx) { + auto *local_md = static_cast(local[desc_idx].metadataP); + auto *remote_md = static_cast(remote[desc_idx].metadataP); + if (!local_md || !remote_md || !remote_md->conn_) { + NIXL_ERROR << "Invalid metadata pointers for descriptor " << desc_idx; + return NIXL_ERR_INVALID_PARAM; + } + + // Validate connection for this descriptor + if (remote_md->conn_ != conn_it->second) { + NIXL_ERROR << "Connection mismatch for descriptor " << desc_idx; + return NIXL_ERR_MISMATCH; + } + // Get transfer info for THIS descriptor + void *transfer_addr = (void *)local[desc_idx].addr; + size_t transfer_size = local[desc_idx].len; + int gpu_id = local[desc_idx].devId; + + NIXL_DEBUG << "Processing descriptor " << desc_idx << " GPU " << gpu_id + << " addr: " << transfer_addr << " size: " << transfer_size; + + NIXL_DEBUG << "DEBUG: remote_agent='" << remote_agent << "' localAgent='" << localAgent + << "'"; + + // Check for same-agent (local) transfer - handle with direct memcpy + if (remote_agent == localAgent) { + NIXL_DEBUG << "Same-agent transfer detected from localAgent= " << localAgent + << "to remote_agent " << remote_agent << "for descriptor " << desc_idx + << ", using memcpy fallback for " << transfer_size << " bytes"; + + // For same-agent transfers, we need to copy directly between the descriptor addresses + // The remote[desc_idx].addr should be the target address for the transfer + void *remote_addr = reinterpret_cast(remote[desc_idx].addr); + + NIXL_DEBUG << "About to perform memcpy: local_addr=" << transfer_addr + << " remote_addr=" << remote_addr << " size=" << transfer_size; + + if (op_type == nixlLibfabricReq::WRITE) { + // Write: copy from local_addr to remote_addr + std::memcpy(remote_addr, transfer_addr, transfer_size); + NIXL_DEBUG << "Same-agent memcpy write completed: " << transfer_addr << " -> " + << remote_addr << " (" << transfer_size << " bytes)"; + } else { + // Read: copy from remote_addr to local_addr + std::memcpy(transfer_addr, remote_addr, transfer_size); + NIXL_DEBUG << "Same-agent memcpy read completed: " << remote_addr << " -> " + << transfer_addr << " (" << transfer_size << " bytes)"; + } + + NIXL_DEBUG << "Successfully processed same-agent descriptor " << desc_idx + << " using memcpy fallback"; + continue; // Skip the rail manager transfer for this descriptor + } + + // Prepare and submit transfer for remote agents + nixl_status_t status = rail_manager.prepareAndSubmitTransfer( + op_type, + transfer_addr, + transfer_size, + remote_md->remote_buf_addr_, + local_md->selected_rails_, + local_md->rail_mr_list_, + remote_md->rail_remote_key_list_, + conn_it->second->rail_remote_addr_list_, + conn_it->second->agent_index_, + [backend_handle]() { + backend_handle->increment_completed_requests(); + }, // Completion callback + binary_notif // Populate BinaryNotification + ); + + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "prepareAndSubmitTransfer failed for descriptor " << desc_idx << " GPU " + << gpu_id; + return status; + } + + NIXL_DEBUG << "Successfully processed descriptor " << desc_idx << " with " + << binary_notif->xfer_id_count << " requests submitted"; + } + + NIXL_DEBUG << "Processing complete: submitted " << binary_notif->xfer_id_count + << " requests from " << desc_count << " descriptors" << " with " + << binary_notif->xfer_id_count << " total XFER_IDs"; + + // For same-agent transfers, we need to set the total to 0 since we bypassed all rail operations + if (remote_agent == localAgent) { + backend_handle->adjust_total_requests(0); + NIXL_DEBUG << "Same-agent transfer: adjusted total requests to 0 (all handled via memcpy)"; + } else { + // Adjust to actual request count after all submissions complete + backend_handle->adjust_total_requests(binary_notif->xfer_id_count); + } + + // Send notification immediately after successful request submission + if (opt_args && opt_args->hasNotif) { + NIXL_DEBUG << "Sending immediate notification after successful request submission"; + + // Set agent name and message in the BinaryNotification + binary_notif->setAgentName(localAgent); + binary_notif->setMessage(opt_args->notifMsg); + + nixl_status_t notif_status = notifSendPriv(remote_agent, control_request); + if (notif_status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to send immediate notification"; + return notif_status; + } + NIXL_DEBUG << "Immediate notification sent successfully with " + << binary_notif->xfer_id_count << " XFER_IDs"; + } + + // Progress data rails to kick off transfers + if (!progress_thread_enabled_) { + nixl_status_t progress_status = rail_manager.progressActiveDataRails(); + if (progress_status == NIXL_IN_PROG) { + return NIXL_IN_PROG; + } + } + + // For very small transfers we can check for local completions immediately. + if (backend_handle->is_completed()) { + return NIXL_SUCCESS; + } + + return NIXL_IN_PROG; +} + +nixl_status_t +nixlLibfabricEngine::checkXfer(nixlBackendReqH *handle) const { + auto backend_handle = static_cast(handle); + + if (!progress_thread_enabled_) { + nixl_status_t progress_status = rail_manager.progressActiveDataRails(); + if (progress_status != NIXL_SUCCESS && progress_status != NIXL_IN_PROG) { + NIXL_ERROR << "Failed to progress data rails in checkXfer"; + return progress_status; + } + } + // Then check for completions after processing any pending completions + if (backend_handle->is_completed()) { + NIXL_DEBUG << "Data transfer completed successfully"; + return NIXL_SUCCESS; + } + return NIXL_IN_PROG; +} + +nixl_status_t +nixlLibfabricEngine::releaseReqH(nixlBackendReqH *handle) const { + // Add any necessary cleanup for libfabric specific request handling + // For example, if we're using a custom request structure: + // nixlLibfabricReqH* req = static_cast(handle); + // // Perform any necessary cleanup + // delete req; + + if (!handle) { + return NIXL_SUCCESS; + } + + // Let NIXL framework handle the deletion + NIXL_DEBUG << "releaseReqH completed successfully"; + return NIXL_SUCCESS; +} + +// notifSendPriv that accept control request +nixl_status_t +nixlLibfabricEngine::notifSendPriv(const std::string &remote_agent, + nixlLibfabricReq *control_request) const { + auto it = connections_.find(remote_agent); + if (it == connections_.end()) { + NIXL_ERROR << "No connection found for agent: " << remote_agent; + return NIXL_ERR_NOT_FOUND; + } + + auto connection = it->second; + const size_t control_rail_id = 0; // Only use control rail 0 for notifications + + // Set the correct buffer size for the notification + control_request->buffer_size = sizeof(BinaryNotification); + + // Get BinaryNotification from control request buffer for logging + BinaryNotification *binary_notif = + reinterpret_cast(control_request->buffer); + + NIXL_DEBUG << "Sending binary notification control request" + << " Message: " << binary_notif->getMessage() + << " xfer_id_count: " << binary_notif->xfer_id_count; + nixl_status_t status = + rail_manager.postControlMessage(nixlLibfabricRailManager::ControlMessageType::NOTIFICATION, + control_request, + connection->control_rail_remote_addr_list_[control_rail_id], + connection->agent_index_); + + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "postControlMessage failed on control rail " << control_rail_id; + return NIXL_ERR_BACKEND; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricEngine::genNotif(const std::string &remote_agent, const std::string &msg) const { + // For regular notifications, we need to allocate a temporary control request + // since we don't have a pre-allocated one from prepXfer + const size_t control_rail_id = 0; + nixlLibfabricReq *control_req = rail_manager.getControlRail(control_rail_id) + .allocateControlRequest(sizeof(BinaryNotification)); + if (!control_req) { + NIXL_ERROR << "Failed to allocate temporary control request for genNotif"; + return NIXL_ERR_BACKEND; + } + + // Create BinaryNotification directly in the control buffer + BinaryNotification *binary_notif = reinterpret_cast(control_req->buffer); + binary_notif->clear(); + binary_notif->setAgentName(localAgent); + binary_notif->setMessage(msg); + + return notifSendPriv(remote_agent, control_req); +} + +nixl_status_t +nixlLibfabricEngine::getNotifs(notif_list_t ¬if_list) { + if (!progress_thread_enabled_) { + nixl_status_t progress_status = rail_manager.progressActiveDataRails(); + if (progress_status != NIXL_SUCCESS && progress_status != NIXL_IN_PROG) { + NIXL_ERROR << "Failed to progress data rails in getNotifs"; + return progress_status; + } + } + + // Then check for available notifications after processing completions + // Thread-safe access to internal notification list + { + std::lock_guard lock(notif_mutex_); + + // Move all notifications from internal list to user's list + notif_list.insert(notif_list.end(), notifMainList_.begin(), notifMainList_.end()); + + if (!notifMainList_.empty()) { + NIXL_DEBUG << "Retrieved " << notifMainList_.size() << " notifications"; + // Clear the internal list after copying + notifMainList_.clear(); + return NIXL_SUCCESS; + } + + // Clear the internal list after copying (even if empty) + notifMainList_.clear(); + } + + return NIXL_IN_PROG; +} + +/**************************************** + * ConnectionManagement Thread Function + *****************************************/ + +// Background progress function that continuously processes completions on all rails +nixl_status_t +nixlLibfabricEngine::cmThread() { + NIXL_DEBUG << "ConnectionManagement thread started successfully"; + NIXL_DEBUG << "Initial receives already posted in main thread, entering progress loop"; + + // Main progress loop - continuously process completions on all rails + while (!cm_thread_stop_.load()) { + + nixl_status_t status = rail_manager.progressAllControlRails(); + if (status == NIXL_SUCCESS) { + NIXL_DEBUG << "Processed completions on control rails"; + } else if (status != NIXL_IN_PROG && status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to process completions on control rails"; + return NIXL_ERR_BACKEND; + } + // Sleep briefly to avoid spinning too aggressively when blocking cq read is not used + if (!rail_manager.getControlRail(0).blocking_cq_sread_supported) { + std::this_thread::sleep_for(std::chrono::nanoseconds(10)); + } + } + NIXL_DEBUG << "ConnectionManagement thread exiting cleanly"; + return NIXL_SUCCESS; +} + +/**************************************** + * Progress Thread Function (Data Rails Only) + *****************************************/ + +// Progress thread that continuously processes completions only on data rails +nixl_status_t +nixlLibfabricEngine::progressThread() { + NIXL_DEBUG << "Progress thread started successfully for data rails only"; + // Main progress loop - continuously process completions only on data rails + while (!progress_thread_stop_.load()) { + // Process completions only on data rails (non-blocking) + bool any_completions = false; + nixl_status_t status = rail_manager.progressActiveDataRails(); + if (status == NIXL_SUCCESS) { + any_completions = true; + NIXL_DEBUG << "Processed completions on data rails"; + } else if (status != NIXL_IN_PROG && status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to process completions on data rails"; + // Don't return error, continue for robustness + } + if (!any_completions) { + std::this_thread::sleep_for(progress_thread_delay_); + } + } + NIXL_DEBUG << "Progress thread exiting cleanly"; + return NIXL_SUCCESS; +} + +void +nixlLibfabricEngine::postShutdownCompletion() { + NIXL_DEBUG << "Posting shutdown signal to wake up background thread"; + // Send shutdown message to self on rail 0 if self-connection exists + auto self_conn_it = connections_.find(localAgent); + if (self_conn_it != connections_.end() && self_conn_it->second && + rail_manager.getNumDataRails() > 0) { + const size_t rail_id = 0; // Use rail 0 for shutdown signal + + // Allocate control request + const size_t control_rail_id = 0; + const size_t shutdown_msg_len = 8; // "SHUTDOWN" length + nixlLibfabricReq *control_request = + rail_manager.getControlRail(control_rail_id).allocateControlRequest(shutdown_msg_len); + if (!control_request) { + NIXL_ERROR << "Failed to allocate control request for shutdown"; + return; + } + + // Copy shutdown message to the control request buffer + std::strcpy(static_cast(control_request->buffer), "SHUTDOWN"); + control_request->buffer_size = shutdown_msg_len; + + nixl_status_t status = rail_manager.postControlMessage( + nixlLibfabricRailManager::ControlMessageType::DISCONNECT_REQ, + control_request, + self_conn_it->second->rail_remote_addr_list_[rail_id], + self_conn_it->second->agent_index_); + + if (status == NIXL_SUCCESS) { + NIXL_DEBUG << "Shutdown signal posted successfully on rail " << rail_id; + } else { + NIXL_ERROR << "Failed to post shutdown signal on rail " << rail_id; + } + } else { + NIXL_ERROR << "Could not find self-connection or rails not initialized"; + } +} + +/**************************************** + * Static Callback Functions + *****************************************/ + +void +nixlLibfabricEngine::processNotification(const std::string &serialized_notif) { + // Only handle binary notification format + // Check if this is a binary notification (fixed size) + NIXL_DEBUG << "Received notification size: " << serialized_notif.size() + << ", sizeof(BinaryNotification): " << sizeof(BinaryNotification); + + if (serialized_notif.size() != sizeof(BinaryNotification)) { + NIXL_ERROR << "Invalid notification size: " << serialized_notif.size() + << ", expected: " << sizeof(BinaryNotification); + return; + } + + // Process binary notification format + const BinaryNotification *binary_notif = + reinterpret_cast(serialized_notif.data()); + + std::string remote_name = binary_notif->getAgentName(); + std::string msg = binary_notif->getMessage(); + std::unordered_set expected_xfer_ids = binary_notif->getXferIds(); + + NIXL_TRACE << "Received binary notification from " << remote_name << " msg: " << msg + << " xfer_id_count: " << binary_notif->xfer_id_count; + + // Check if this is a transfer notification that needs queuing + if (!expected_xfer_ids.empty()) { + std::stringstream xfer_ids_log; + xfer_ids_log << "Expected XFER_IDs from binary notification: ["; + bool first = true; + for (uint32_t xfer_id : expected_xfer_ids) { + if (!first) xfer_ids_log << ", "; + xfer_ids_log << xfer_id; + first = false; + } + xfer_ids_log << "] (total: " << expected_xfer_ids.size() << ")"; + NIXL_TRACE << xfer_ids_log.str(); + // Check if all expected XFER_IDs have already arrived + if (allXferIdsReceived(expected_xfer_ids)) { + NIXL_TRACE + << "All XFER_IDs already received, processing binary notification immediately"; + std::lock_guard lock(notif_mutex_); + notifMainList_.push_back({remote_name, msg}); + NIXL_DEBUG << "Binary notification processed immediately: " << msg; + } else { + NIXL_TRACE << "Not all XFER_IDs received yet, queuing binary notification"; + std::lock_guard lock(receiver_tracking_mutex_); + pending_notifications_.emplace_back(remote_name, msg, expected_xfer_ids); + NIXL_TRACE << "Binary notification queued for later processing: " << msg; + } + } else { + // Regular notification without XFER_IDs - process immediately + NIXL_TRACE << "Regular binary notification (no XFER_IDs), processing immediately"; + std::lock_guard lock(notif_mutex_); + notifMainList_.push_back({remote_name, msg}); + NIXL_TRACE << "Regular binary notification processed immediately: " << msg; + } +} + +void +nixlLibfabricEngine::processConnectionAck(uint16_t agent_idx, + nixlLibfabricConnection *conn_info, + ConnectionState state) { + std::string remote_agent_name = agent_names_[agent_idx]; + NIXL_DEBUG << "Connection state callback for agent " << remote_agent_name + << " agent_idx: " << agent_idx; + std::lock_guard lock(connections_[remote_agent_name]->conn_state_mutex_); + connections_[remote_agent_name]->overall_state_ = ConnectionState::CONNECTED; + connections_[remote_agent_name]->cv_.notify_all(); + NIXL_DEBUG << "Connection state updated to CONNECTED"; +} + +nixl_status_t +nixlLibfabricEngine::processConnectionRequest(uint16_t agent_idx, + const std::string &serialized_data, + nixlLibfabricRail *rail) { + NIXL_DEBUG << "Processing connection request from agent " << agent_idx << " on rail " + << rail->rail_id; + + // Use rail manager to deserialize ALL endpoints at once with "src" prefix (connection request + // contains source endpoints) + std::vector> data_endpoints; + std::vector> control_endpoints; + nixl_status_t status = rail_manager.deserializeConnectionInfo( + "src", serialized_data, data_endpoints, control_endpoints); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to deserialize connection info"; + return status; + } + + // Insert ALL data rail addresses at once + std::vector data_fi_addrs; + std::vector data_ep_names; + status = rail_manager.insertAllAddresses( + nixlLibfabricRailManager::RailType::DATA, data_endpoints, data_fi_addrs, data_ep_names); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to insert data rail addresses"; + return status; + } + + // Insert ALL control rail addresses at once + std::vector control_fi_addrs; + std::vector control_ep_names; + status = rail_manager.insertAllAddresses(nixlLibfabricRailManager::RailType::CONTROL, + control_endpoints, + control_fi_addrs, + control_ep_names); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to insert control rail addresses"; + return status; + } + + // Use the first control rail's fi_addr for ACK (same as before) + fi_addr_t initiator_control_fi_addr = control_fi_addrs[0]; + + NIXL_DEBUG << "Successfully inserted addresses for " << data_fi_addrs.size() + << " data rails and " << control_fi_addrs.size() << " control rails" + << ", initiator_control_fi_addr: " << initiator_control_fi_addr; + + // Send acknowledgement back to the initiator using the rail manager + size_t ep_name_len = sizeof(rail->ep_name); + + // Allocate control request + const size_t control_rail_id = 0; + nixlLibfabricReq *control_request = + rail_manager.getControlRail(control_rail_id).allocateControlRequest(ep_name_len); + if (!control_request) { + NIXL_ERROR << "Failed to allocate control request for connection ACK"; + return NIXL_ERR_BACKEND; + } + + // Copy endpoint name to control request buffer + std::memcpy(control_request->buffer, rail->ep_name, ep_name_len); + control_request->buffer_size = ep_name_len; + + nixl_status_t ack_status = rail_manager.postControlMessage( + nixlLibfabricRailManager::ControlMessageType::CONNECTION_ACK, + control_request, + initiator_control_fi_addr, + agent_idx); + if (ack_status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to send ACK via rail manager"; + return ack_status; + } + + NIXL_DEBUG << "ACK sent successfully via rail manager"; + return NIXL_SUCCESS; +} + +/**************************************** + * Receiver Side XFER_ID Tracking Helper Methods + *****************************************/ + +void +nixlLibfabricEngine::addReceivedXferId(uint32_t xfer_id) { + { + std::lock_guard lock(receiver_tracking_mutex_); + received_remote_writes_.insert(xfer_id); + NIXL_DEBUG << "Added received XFER_ID " << xfer_id + << " to global tracking set (total: " << received_remote_writes_.size() << ")"; + } + + // Check if any pending notifications can now be processed + checkPendingNotifications(); +} + +bool +nixlLibfabricEngine::allXferIdsReceived(const std::unordered_set &expected) { + std::lock_guard lock(receiver_tracking_mutex_); + // Check if all expected XFER_IDs are in the received set + for (uint32_t xfer_id : expected) { + if (received_remote_writes_.find(xfer_id) == received_remote_writes_.end()) { + NIXL_TRACE << "XFER_ID " << xfer_id << " not yet received"; + return false; + } + } + NIXL_DEBUG << "All " << expected.size() << " expected XFER_IDs have been received"; + return true; +} + +/**************************************** + * Notification Queuing Helper Methods + *****************************************/ + +void +nixlLibfabricEngine::checkPendingNotifications() { + std::lock_guard lock(receiver_tracking_mutex_); + auto it = pending_notifications_.begin(); + while (it != pending_notifications_.end()) { + // Check if all expected XFER_IDs for this notification have arrived + bool all_received = true; + for (uint32_t xfer_id : it->expected_xfer_ids) { + if (received_remote_writes_.find(xfer_id) == received_remote_writes_.end()) { + all_received = false; + break; + } + } + + if (all_received) { + NIXL_TRACE << "All XFER_IDs received for queued notification, processing now"; + + // Move notification to main list (need to acquire notif_mutex_) + { + std::lock_guard notif_lock(notif_mutex_); + notifMainList_.push_back({it->remote_agent, it->message}); + } + + NIXL_TRACE << "Processed queued notification: " << it->message; + + // Remove from pending list + it = pending_notifications_.erase(it); + } else { + ++it; + } + } +} + +void +nixlLibfabricEngine::cleanup() { + NIXL_DEBUG << "Cleaning up all resources"; +#ifdef HAVE_CUDA + // Cleanup CUDA context + vramFiniCtx(); +#endif + + NIXL_DEBUG << "Cleanup all resources complete"; +} diff --git a/src/plugins/libfabric/libfabric_backend.h b/src/plugins/libfabric/libfabric_backend.h new file mode 100644 index 0000000000..0b07f4c2ef --- /dev/null +++ b/src/plugins/libfabric/libfabric_backend.h @@ -0,0 +1,584 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_PLUGINS_LIBFABRIC_LIBFABRIC_BACKEND_H +#define NIXL_SRC_PLUGINS_LIBFABRIC_LIBFABRIC_BACKEND_H + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "nixl.h" +#include "backend/backend_engine.h" +#include "common/nixl_time.h" +#include "serdes/serdes.h" + +#include "libfabric/libfabric_rail_manager.h" +#include "libfabric/libfabric_common.h" + +#ifdef HAVE_CUDA +#include +#include +#endif + +// Forward declarations +class nixlLibfabricEngine; + +#ifdef HAVE_CUDA +/** CUDA context management for libfabric backend */ +class nixlLibfabricCudaCtx { +private: + CUcontext pthrCudaCtx_; + int myDevId_; + +public: + nixlLibfabricCudaCtx() { + pthrCudaCtx_ = NULL; + myDevId_ = -1; + } + + /** Reset CUDA context pointer to initial state */ + void + cudaResetCtxPtr(); + + /** Update CUDA context pointer for given memory address and device */ + int + cudaUpdateCtxPtr(void *address, int expected_dev, bool &was_updated); + + /** Set the current CUDA context */ + int + cudaSetCtx(); +}; +#endif + +/** Private metadata for locally registered memory */ +class nixlLibfabricPrivateMetadata : public nixlBackendMD { +private: + void *buffer_; // Local memory buffer address + size_t length_; // Buffer length in bytes + int gpu_device_id_; // GPU device ID for VRAM, -1 for DRAM + std::vector rail_mr_list_; // Memory registrations, one per rail + std::vector rail_key_list_; // Remote access keys, one per rail + std::vector src_ep_names_; // Source endpoint names, one per rail + std::vector selected_rails_; // Rails selected based on memory topology + +public: + nixlLibfabricPrivateMetadata() : nixlBackendMD(true), gpu_device_id_(-1) {} + friend class nixlLibfabricEngine; +}; + +/** Public metadata for remote memory access */ +class nixlLibfabricPublicMetadata : public nixlBackendMD { +private: + uint64_t remote_buf_addr_; // Remote buffer base address + std::shared_ptr conn_; // Connection to remote agent + std::vector rail_remote_key_list_; // Remote access keys, one per rail + std::vector src_ep_names_; // Source endpoint names, one per rail + +public: + nixlLibfabricPublicMetadata() : nixlBackendMD(false) {} + friend class nixlLibfabricEngine; +}; + +/** Multi-rail connection metadata for remote agents */ +class nixlLibfabricConnection : public nixlBackendConnMD { +private: + size_t agent_index_; // Unique agent identifier in agent_names vector + std::string remoteAgent_; // Remote agent name + std::vector rail_remote_addr_list_; // Data rail libfabric addresses + std::vector control_rail_remote_addr_list_; // Control rail libfabric addresses + std::vector src_ep_names_; // Data rail endpoint names + std::vector control_ep_names_; // Control rail endpoint names + ConnectionState overall_state_; // Current connection state + std::mutex conn_state_mutex_; // Protects connection state + std::condition_variable cv_; // For blocking connection establishment + size_t num_connected_rails_; // Number of successfully connected rails + std::string initiator_addr_; // Local endpoint address + std::string remote_addr_; // Remote endpoint address +public: + friend class nixlLibfabricEngine; + friend class nixlLibfabricRail; +}; + +/** Request handle for multi-rail transfer operations */ +class nixlLibfabricBackendH : public nixlBackendReqH { +private: + std::atomic completed_requests_; // Atomic count of completed requests + std::atomic total_requests_used_; // Total number of requests for this transfer + +public: + nixlLibfabricBackendH(); + ~nixlLibfabricBackendH(); + + /** Check if all requests in this transfer have completed */ + bool + is_completed() const; + + /** Initialize completion tracking for multi-request transfer */ + void + init_request_tracking(size_t num_requests); + + /** Atomically increment completed request count */ + void + increment_completed_requests(); + + /** Get current count of completed requests */ + size_t + get_completed_requests_count() const; + + /** Get total number of requests used for this transfer */ + size_t + get_total_requests_used() const; + + /** Adjust total request count to actual value after submissions complete */ + void + adjust_total_requests(size_t actual_count); +}; + +class nixlLibfabricEngine : public nixlBackendEngine { + friend class nixlLibfabricRail; // Allow nixlLibfabricRail to access private members + +private: + // Threading infrastructure - declared first to match initialization order + std::atomic cm_thread_stop_; + + // Store user's original progress thread preference + bool progress_thread_enabled_; + + // Progress thread delay in microseconds + std::chrono::microseconds progress_thread_delay_; + + // Rail Manager - Stack allocated for better performance (mutable for const methods) + mutable nixlLibfabricRailManager rail_manager; + + // Configurable striping threshold + size_t striping_threshold_; + + mutable size_t total_transfer_size_; + + // Map of agent name to connection info + // > + mutable std::unordered_map> connections_; + mutable std::vector agent_names_; // List of agent names for easy access + + // Threading infrastructure - remaining members + // Connection Management (CM) thread + std::mutex cm_mutex_; + std::thread cm_thread_; + std::condition_variable cm_cv_; + + // Progress thread for data rail CQs only + std::thread progress_thread_; + std::atomic progress_thread_stop_; + + // Mutex for connection state tracking + mutable std::mutex connection_state_mutex_; + + void + cleanup(); + + // Central notification storage + std::mutex notif_mutex_; + notif_list_t notifMainList_; + + // Receiver Side XFER_ID Tracking + std::mutex receiver_tracking_mutex_; + std::unordered_set received_remote_writes_; // All received XFER_IDs (global) + + // Notification Queuing + struct PendingNotification { + std::string remote_agent; + std::string message; + std::unordered_set expected_xfer_ids; // From ref_xfer_id_list + std::chrono::steady_clock::time_point received_time; + + PendingNotification(const std::string &agent, + const std::string &msg, + const std::unordered_set &xfer_ids) + : remote_agent(agent), + message(msg), + expected_xfer_ids(xfer_ids), + received_time(std::chrono::steady_clock::now()) {} + }; + + std::vector pending_notifications_; + + // Connection management helpers + nixl_status_t + establishConnection(const std::string &remote_agent) const; + + // Common connection creation helper + nixl_status_t + createAgentConnection(const std::string &agent_name, + const std::vector> &data_rail_endpoints, + const std::vector> &control_rail_endpoints); + // Private notification implementation with unified binary notification system + nixl_status_t + notifSendPriv(const std::string &remote_agent, nixlLibfabricReq *control_request) const; +#ifdef HAVE_CUDA + // CUDA context management + std::unique_ptr cudaCtx_; + bool cuda_addr_wa_; // CUDA address workaround flag +#endif + + // ConnectionManagement thread and completion processing + nixl_status_t + cmThread(); + void + postShutdownCompletion(); + // Progress thread for data rail CQs only + nixl_status_t + progressThread(); + + + // Engine message processing methods + void + processNotification(const std::string &serialized_notif); + void + processConnectionAck(uint16_t agent_idx, + nixlLibfabricConnection *conn_info, + ConnectionState state); + nixl_status_t + processConnectionRequest(uint16_t agent_idx, + const std::string &serialized_data, + nixlLibfabricRail *rail); + + +#ifdef HAVE_CUDA + // CUDA context management methods + void + vramInitCtx(); + int + vramUpdateCtx(void *address, uint64_t devId, bool &restart_reqd); + int + vramApplyCtx(); + void + vramFiniCtx(); +#endif + +public: + /** Initialize multi-rail libfabric backend engine */ + nixlLibfabricEngine(const nixlBackendInitParams *init_params); + /** Destroy engine and cleanup all resources */ + ~nixlLibfabricEngine(); + + bool + supportsRemote() const override { + return true; + } + + bool + supportsLocal() const override { + return true; + } + + bool + supportsNotif() const override { + return true; + } + + /** Get list of supported memory types */ + nixl_mem_list_t + getSupportedMems() const override; + + /* Object management */ + /** Serialize memory metadata for remote access */ + nixl_status_t + getPublicData(const nixlBackendMD *meta, std::string &str) const override; + + /** Get local connection information for all rails */ + nixl_status_t + getConnInfo(std::string &str) const override; + + /** Load remote agent connection information */ + nixl_status_t + loadRemoteConnInfo(const std::string &remote_agent, + const std::string &remote_conn_info) override; + + /** Establish connection to remote agent */ + nixl_status_t + connect(const std::string &remote_agent) override; + /** + * @brief Gracefully disconnects from a remote agent and cleans up associated resources + * + * This function performs a complete disconnect sequence that ensures proper cleanup + * of all libfabric resources and notifies the remote peer of the disconnection. + * + * The disconnect process follows these steps: + * 1. Validates that an active connection exists for the specified remote agent + * 2. Sends disconnect notification message to the remote peer via control rail, with best + * effort semantics + * 3. Releases libfabric resources like address vector entries via rail manager + * 4. Updates internal connection state to DISCONNECTED + * 5. Removes connection entry from the active connection map + * + * @param[in] remote_agent The identifier of the remote agent to disconnect from + * + * @return nixl_status_t Status code indicating the result of the disconnect operation + * @retval NIXL_SUCCESS Connection successfully disconnected and cleaned up + * @retval NIXL_ERR_NOT_FOUND No active connection exists for the specified remote agent + * @retval NIXL_ERR_TIMEOUT Remote peer did not acknowledge disconnect within timeout period + * @retval NIXL_ERR_BACKEND Libfabric resource cleanup failed + * + * @note This function is thread-safe and can be called concurrently from multiple threads + * @note Active transfers will be cancelled, which may result in incomplete data transmission + * @note The function implements best-effort graceful disconnect; if acknowledgment times out, + * local cleanup still proceeds to prevent resource leaks + * + * @warning Calling disconnect on a non-existent connection returns NIXL_ERR_NOT_FOUND + * @warning This operation is irreversible; a new connection must be established to resume + * communication + * + * @see connect() for establishing connections + * @see getConnectionState() for querying connection status + */ + nixl_status_t + disconnect(const std::string &remote_agent) override; + + /** + * @brief Register memory for RDMA operations with GPU Direct support + * + * Registers memory buffer with libfabric on topology-appropriate rails. + * Supports both DRAM and VRAM with automatic rail selection based on + * memory location and system topology. + * + * @param[in] mem Memory descriptor with address, length, and device info + * @param[in] nixl_mem Memory type (DRAM_SEG or VRAM_SEG) + * @param[out] out Private metadata containing registration information + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + registerMem(const nixlBlobDesc &mem, const nixl_mem_t &nixl_mem, nixlBackendMD *&out) override; + + /** + * @brief Deregister memory from libfabric + * + * Cleans up memory registrations on all rails where the memory was registered. + * + * @param[in] meta Private metadata from registerMem() + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + deregisterMem(nixlBackendMD *meta) override; + + /** + * @brief Create public metadata from local private metadata + * + * Converts private memory registration into public metadata that can be + * used for local transfers (loopback operations). + * + * @param[in] input Private metadata from registerMem() + * @param[out] output Public metadata for local operations + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + loadLocalMD(nixlBackendMD *input, nixlBackendMD *&output) override; + + /** + * @brief Create public metadata from remote serialized data + * + * Deserializes remote memory information and creates public metadata + * for accessing remote memory via RDMA operations. + * + * @param[in] input Blob descriptor containing serialized remote metadata + * @param[in] nixl_mem Memory type of the remote memory + * @param[in] remote_agent Name of the remote agent owning the memory + * @param[out] output Public metadata for remote memory access + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + loadRemoteMD(const nixlBlobDesc &input, + const nixl_mem_t &nixl_mem, + const std::string &remote_agent, + nixlBackendMD *&output) override; + + /** + * @brief Release metadata resources + * + * Cleans up metadata object and associated resources. + * + * @param[in] input Metadata to release + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + unloadMD(nixlBackendMD *input) override; + + // Data transfer + /** + * @brief Prepare transfer operation handle + * + * Creates and initializes a request handle for upcoming data transfer. + * Validates connection and prepares internal structures. + * + * @param[in] operation Transfer operation type (NIXL_WRITE or NIXL_READ) + * @param[in] local Local memory descriptors + * @param[in] remote Remote memory descriptors + * @param[in] remote_agent Target remote agent name + * @param[out] handle Request handle for the transfer + * @param[in] opt_args Optional arguments (notifications, etc.) + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const override; + + /** + * @brief Estimate transfer cost and timing + * + * Provides cost estimation for the specified transfer operation. + * Currently returns success without detailed estimation. + * + * @param[in] operation Transfer operation type + * @param[in] local Local memory descriptors + * @param[in] remote Remote memory descriptors + * @param[in] remote_agent Target remote agent name + * @param[in] handle Request handle + * @param[out] duration Estimated transfer duration + * @param[out] err_margin Error margin for the estimate + * @param[out] method Cost estimation method used + * @param[in] opt_args Optional arguments + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + estimateXferCost(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *const &handle, + std::chrono::microseconds &duration, + std::chrono::microseconds &err_margin, + nixl_cost_t &method, + const nixl_opt_args_t *opt_args = nullptr) const override; + + /** + * @brief Execute data transfer with multi-rail striping + * + * Performs high-performance data transfer using multiple rails with automatic + * striping for large transfers or round-robin for small transfers. Supports + * immediate notifications and completion tracking. + * + * @param[in] operation Transfer operation type (NIXL_WRITE or NIXL_READ) + * @param[in] local Local memory descriptors + * @param[in] remote Remote memory descriptors + * @param[in] remote_agent Target remote agent name + * @param[in,out] handle Request handle for tracking completion + * @param[in] opt_args Optional arguments including notifications + * @return NIXL_SUCCESS if complete, NIXL_IN_PROG if ongoing, error code on failure + */ + nixl_status_t + postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const override; + + /** + * @brief Check transfer completion status + * + * Polls for transfer completion and processes any pending completions + * if progress thread is disabled. + * + * @param[in] handle Request handle to check + * @return NIXL_SUCCESS if complete, NIXL_IN_PROG if ongoing, error code on failure + */ + nixl_status_t + checkXfer(nixlBackendReqH *handle) const override; + + /** + * @brief Release request handle resources + * + * Cleans up request handle after transfer completion. + * + * @param[in] handle Request handle to release + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + releaseReqH(nixlBackendReqH *handle) const override; + + // Notification system + /** + * @brief Retrieve available notifications + * + * Returns all currently available notifications and processes any pending + * completions if progress thread is disabled. + * + * @param[out] notif_list List to store retrieved notifications + * @return NIXL_SUCCESS if notifications available, NIXL_IN_PROG if none, error code on failure + */ + nixl_status_t + getNotifs(notif_list_t ¬if_list) override; + + /** + * @brief Send notification to remote agent + * + * Sends a notification message to the specified remote agent using + * the binary notification protocol. + * + * @param[in] remote_agent Target remote agent name + * @param[in] msg Notification message to send + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + genNotif(const std::string &remote_agent, const std::string &msg) const override; + + // Receiver Side XFER_ID Tracking Helper Methods + /** + * @brief Add received XFER_ID to global tracking set + * + * Thread-safe method to track received data transfers and trigger + * processing of pending notifications when all expected transfers arrive. + * + * @param[in] xfer_id Transfer ID that was received + */ + void + addReceivedXferId(uint32_t xfer_id); + + /** + * @brief Check if all expected XFER_IDs have been received + * + * Determines if all transfers associated with a notification have completed. + * + * @param[in] expected Set of expected XFER_IDs + * @return true if all expected IDs received, false otherwise + */ + bool + allXferIdsReceived(const std::unordered_set &expected); + + // Notification Queuing Helper Methods + /** + * @brief Process pending notifications that are now ready + * + * Checks pending notifications to see if their associated transfers + * have completed and moves them to the main notification list. + */ + void + checkPendingNotifications(); +}; + +#endif diff --git a/src/plugins/libfabric/libfabric_plugin.cpp b/src/plugins/libfabric/libfabric_plugin.cpp new file mode 100644 index 0000000000..096948956a --- /dev/null +++ b/src/plugins/libfabric/libfabric_plugin.cpp @@ -0,0 +1,40 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "backend/backend_plugin.h" +#include "libfabric_backend.h" + +// Plugin type alias for convenience +using libfabric_plugin_t = nixlBackendPluginCreator; + +#ifdef STATIC_PLUGIN_LIBFABRIC +nixlBackendPlugin * +createStaticLIBFABRICPlugin() { + return libfabric_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "LIBFABRIC", "0.1.0", {}, {DRAM_SEG, VRAM_SEG}); +} +#else +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return libfabric_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "LIBFABRIC", "0.1.0", {}, {DRAM_SEG, VRAM_SEG}); +} + +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} +#endif diff --git a/src/plugins/libfabric/meson.build b/src/plugins/libfabric/meson.build new file mode 100644 index 0000000000..c48d13806b --- /dev/null +++ b/src/plugins/libfabric/meson.build @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# LibFabric plugin configuration + +# Enable libfabric utils layer +libfabric_plugin_deps = [ + nixl_infra, + nixl_common_dep, + libfabric_dep, + serdes_interface, + libfabric_utils_dep, +] + +compile_flags = ['-DHAVE_LIBFABRIC'] + +# Add CUDA support if available +if cuda_dep.found() + libfabric_plugin_deps += [cuda_dep] + compile_flags += ['-DHAVE_CUDA'] +endif + +# Build as static or shared library based on configuration +if 'LIBFABRIC' in static_plugins + libfabric_backend_lib = static_library( + 'LIBFABRIC', + 'libfabric_backend.cpp', + 'libfabric_backend.h', + 'libfabric_plugin.cpp', + dependencies: libfabric_plugin_deps, + include_directories: [nixl_inc_dirs, utils_inc_dirs], + install: false, + cpp_args: compile_flags, + name_prefix: 'libplugin_', + ) +else + libfabric_backend_lib = shared_library( + 'LIBFABRIC', + 'libfabric_backend.cpp', + 'libfabric_backend.h', + 'libfabric_plugin.cpp', + dependencies: libfabric_plugin_deps, + include_directories: [nixl_inc_dirs, utils_inc_dirs], + install: true, + cpp_args: compile_flags + ['-fPIC'], + name_prefix: 'libplugin_', + install_dir: plugin_install_dir, + ) + + # Add to plugin list for debug builds + if get_option('buildtype') == 'debug' + run_command( + 'sh', + '-c', + 'echo "LIBFABRIC=' + libfabric_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', + check: true, + ) + endif +endif + +# Declare dependency interface +libfabric_backend_interface = declare_dependency(link_with: libfabric_backend_lib) diff --git a/src/plugins/meson.build b/src/plugins/meson.build index cc11b4a6d6..551c4d1ff7 100644 --- a/src/plugins/meson.build +++ b/src/plugins/meson.build @@ -21,6 +21,10 @@ subdir('ucx_mo') subdir('posix') # Always try to build POSIX backend, it will handle its own dependencies subdir('obj') # Always try to build Obj backend, it will handle its own dependencies +if libfabric_dep.found() + subdir('libfabric') +endif + disable_gds_backend = get_option('disable_gds_backend') if not disable_gds_backend and cuda_dep.found() subdir('cuda_gds') diff --git a/src/plugins/mooncake/README.md b/src/plugins/mooncake/README.md index e0fedcecce..8ab961f481 100644 --- a/src/plugins/mooncake/README.md +++ b/src/plugins/mooncake/README.md @@ -26,8 +26,10 @@ Mooncake transfer engine is a high-performance, zero-copy data transfer library. 3. To test the Mooncake backend, you can run the unit test in `test/unit/plugins/mooncake/mooncake_backend_test`. +4. To use the Notify feature, you need to download the latest main branch of Mooncake. + ## Known Issues -1. The `Notif[ication]` and `ProgTh[read]` features are not supported. +1. The `ProgTh[read]` features are not supported. 2. The current version of Mooncake Transfer Engine manages metadata exchange by itself, which is different from NIXL. 3. The sum of the number of release requests for each handle allocated by `prepXfer()` should be less than `kMaxRequestCount(1024)`. diff --git a/src/plugins/mooncake/meson.build b/src/plugins/mooncake/meson.build index 777849ce0e..759aefe736 100644 --- a/src/plugins/mooncake/meson.build +++ b/src/plugins/mooncake/meson.build @@ -21,7 +21,7 @@ compile_flags = [] if 'Mooncake' in static_plugins mooncake_backend_lib = static_library('Mooncake', 'mooncake_backend.cpp', 'mooncake_backend.h', 'mooncake_plugin.cpp', - dependencies: [nixl_infra, serdes_interface, mooncake_lib, cuda_dep], + dependencies: [nixl_infra, nixl_common_dep, serdes_interface, mooncake_lib, cuda_dep], include_directories: nixl_inc_dirs, install: false, cpp_args : compile_flags, @@ -29,12 +29,13 @@ if 'Mooncake' in static_plugins else mooncake_backend_lib = shared_library('Mooncake', 'mooncake_backend.cpp', 'mooncake_backend.h', 'mooncake_plugin.cpp', - dependencies: [nixl_infra, serdes_interface, mooncake_lib, cuda_dep], + dependencies: [nixl_infra, nixl_common_dep, serdes_interface, mooncake_lib, cuda_dep], include_directories: [nixl_inc_dirs, utils_inc_dirs], install: true, cpp_args : compile_flags + ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', diff --git a/src/plugins/mooncake/mooncake_backend.cpp b/src/plugins/mooncake/mooncake_backend.cpp index a3a7d33574..45b61649c1 100644 --- a/src/plugins/mooncake/mooncake_backend.cpp +++ b/src/plugins/mooncake/mooncake_backend.cpp @@ -16,6 +16,7 @@ */ #include "mooncake_backend.h" #include "serdes/serdes.h" +#include "common/nixl_log.h" #include #include @@ -23,8 +24,10 @@ #include #include #include +#include -std::vector findLocalIpAddresses() { +std::vector +findLocalIpAddresses() { std::vector ips; struct ifaddrs *ifaddr, *ifa; @@ -42,9 +45,20 @@ std::vector findLocalIpAddresses() { continue; } + // Check if interface is UP and RUNNING + if (!(ifa->ifa_flags & IFF_UP) || !(ifa->ifa_flags & IFF_RUNNING)) { + NIXL_INFO << "Skipping interface " << ifa->ifa_name << " (not UP or not RUNNING)"; + continue; + } + char host[NI_MAXHOST]; - if (getnameinfo(ifa->ifa_addr, sizeof(struct sockaddr_in), host, - NI_MAXHOST, nullptr, 0, NI_NUMERICHOST) == 0) { + if (getnameinfo(ifa->ifa_addr, + sizeof(struct sockaddr_in), + host, + NI_MAXHOST, + nullptr, + 0, + NI_NUMERICHOST) == 0) { ips.push_back(host); } } @@ -54,20 +68,19 @@ std::vector findLocalIpAddresses() { return ips; } -nixlMooncakeEngine::nixlMooncakeEngine (const nixlBackendInitParams* init_params) -: nixlBackendEngine (init_params) { +nixlMooncakeEngine::nixlMooncakeEngine(const nixlBackendInitParams *init_params) + : nixlBackendEngine(init_params) { local_agent_name_ = init_params->localAgent; auto ips = findLocalIpAddresses(); std::string segment_name = "127.0.0.1"; if (!ips.empty()) segment_name = ips[0]; if (getenv("NIXL_MOONCAKE_IP_ADDR")) segment_name = std::string(getenv("NIXL_MOONCAKE_IP_ADDR")); - engine_ = createTransferEngine("P2PHANDSHAKE", - segment_name.c_str(), - "", 0, true); + engine_ = createTransferEngine("P2PHANDSHAKE", segment_name.c_str(), "", 0, true); } -nixl_mem_list_t nixlMooncakeEngine::getSupportedMems () const { +nixl_mem_list_t +nixlMooncakeEngine::getSupportedMems() const { nixl_mem_list_t mems; mems.push_back(DRAM_SEG); mems.push_back(VRAM_SEG); @@ -75,7 +88,7 @@ nixl_mem_list_t nixlMooncakeEngine::getSupportedMems () const { } // Through parent destructor the unregister will be called. -nixlMooncakeEngine::~nixlMooncakeEngine () { +nixlMooncakeEngine::~nixlMooncakeEngine() { destroyTransferEngine(engine_); } @@ -88,17 +101,20 @@ nixlMooncakeEngine::~nixlMooncakeEngine () { // (segment name in the context of Mooncake Transfer Engine). // loadRemoteConnInfo() opens the segment, which implicitly retrieves metadata // (such as QP numbers) of the remote agent. -nixl_status_t nixlMooncakeEngine::connect(const std::string &remote_agent) { +nixl_status_t +nixlMooncakeEngine::connect(const std::string &remote_agent) { return NIXL_SUCCESS; } // TODO We purposely set this function as empty. // Will be changed to follow NIXL's paradigm after refactoring Mooncake Transfer Engine. -nixl_status_t nixlMooncakeEngine::disconnect(const std::string &remote_agent) { +nixl_status_t +nixlMooncakeEngine::disconnect(const std::string &remote_agent) { return NIXL_SUCCESS; } -nixl_status_t nixlMooncakeEngine::getConnInfo(std::string &str) const { +nixl_status_t +nixlMooncakeEngine::getConnInfo(std::string &str) const { const static size_t kBufLen = 64; char buf_out[kBufLen]; getLocalIpAndPort(engine_, buf_out, kBufLen); @@ -106,28 +122,30 @@ nixl_status_t nixlMooncakeEngine::getConnInfo(std::string &str) const { return NIXL_SUCCESS; } -nixl_status_t nixlMooncakeEngine::loadRemoteConnInfo (const std::string &remote_agent, - const std::string &remote_conn_info) -{ +nixl_status_t +nixlMooncakeEngine::loadRemoteConnInfo(const std::string &remote_agent, + const std::string &remote_conn_info) { std::lock_guard lock(mutex_); auto segment_id = openSegment(engine_, remote_conn_info.c_str()); if (segment_id < 0) return NIXL_ERR_BACKEND; - connected_agents_[remote_agent].segment_id = segment_id; + connected_agents_[remote_agent].segment_id = segment_id; return NIXL_SUCCESS; } struct nixlMooncakeBackendMD : public nixlBackendMD { nixlMooncakeBackendMD(bool isPrivate) : nixlBackendMD(isPrivate) {} - virtual ~nixlMooncakeBackendMD(){} + + virtual ~nixlMooncakeBackendMD() {} + void *addr; size_t length; int ref_cnt; }; -nixl_status_t nixlMooncakeEngine::registerMem (const nixlBlobDesc &mem, - const nixl_mem_t &nixl_mem, - nixlBackendMD* &out) -{ +nixl_status_t +nixlMooncakeEngine::registerMem(const nixlBlobDesc &mem, + const nixl_mem_t &nixl_mem, + nixlBackendMD *&out) { std::lock_guard lock(mutex_); if (mem_reg_info_.count(mem.addr)) { auto priv = mem_reg_info_[mem.addr]; @@ -135,10 +153,10 @@ nixl_status_t nixlMooncakeEngine::registerMem (const nixlBlobDesc &mem, out = priv; return NIXL_SUCCESS; } - int err = registerLocalMemory(engine_, (void *) mem.addr, mem.len, "*", 1); + int err = registerLocalMemory(engine_, (void *)mem.addr, mem.len, "*", 1); if (err) return NIXL_ERR_BACKEND; auto priv = new nixlMooncakeBackendMD(true); - priv->addr = (void *) mem.addr; + priv->addr = (void *)mem.addr; priv->length = mem.len; priv->ref_cnt = 1; out = priv; @@ -146,10 +164,10 @@ nixl_status_t nixlMooncakeEngine::registerMem (const nixlBlobDesc &mem, return NIXL_SUCCESS; } -nixl_status_t nixlMooncakeEngine::deregisterMem (nixlBackendMD* meta) -{ +nixl_status_t +nixlMooncakeEngine::deregisterMem(nixlBackendMD *meta) { std::lock_guard lock(mutex_); - auto priv = (nixlMooncakeBackendMD *) meta; + auto priv = (nixlMooncakeBackendMD *)meta; priv->ref_cnt--; if (priv->ref_cnt) return NIXL_SUCCESS; int err = unregisterLocalMemory(engine_, priv->addr); @@ -164,81 +182,79 @@ nixl_status_t nixlMooncakeEngine::deregisterMem (nixlBackendMD* meta) // Mooncake Transfer Engine exchanges metadata by itself without any explicit interface, // which is different from NIXL's paradigm. // Therefore no metadata needs to be exposed to the outside. -nixl_status_t nixlMooncakeEngine::getPublicData (const nixlBackendMD* meta, - std::string &str) const -{ +nixl_status_t +nixlMooncakeEngine::getPublicData(const nixlBackendMD *meta, std::string &str) const { return NIXL_SUCCESS; } // TODO We purposely set this function as empty. // Will be changed to follow NIXL's paradigm after refactoring Mooncake Transfer Engine. nixl_status_t -nixlMooncakeEngine::loadLocalMD (nixlBackendMD* input, - nixlBackendMD* &output) -{ +nixlMooncakeEngine::loadLocalMD(nixlBackendMD *input, nixlBackendMD *&output) { output = nullptr; return NIXL_SUCCESS; } // TODO We purposely set this function as empty. // Will be changed to follow NIXL's paradigm after refactoring Mooncake Transfer Engine. -nixl_status_t nixlMooncakeEngine::loadRemoteMD (const nixlBlobDesc &input, - const nixl_mem_t &nixl_mem, - const std::string &remote_agent, - nixlBackendMD* &output) -{ +nixl_status_t +nixlMooncakeEngine::loadRemoteMD(const nixlBlobDesc &input, + const nixl_mem_t &nixl_mem, + const std::string &remote_agent, + nixlBackendMD *&output) { output = nullptr; return NIXL_SUCCESS; } // TODO We purposely set this function as empty. // Will be changed to follow NIXL's paradigm after refactoring Mooncake Transfer Engine. -nixl_status_t nixlMooncakeEngine::unloadMD (nixlBackendMD* input) -{ +nixl_status_t +nixlMooncakeEngine::unloadMD(nixlBackendMD *input) { return NIXL_SUCCESS; } struct nixlMooncakeBackendReqH : public nixlBackendReqH { nixlMooncakeBackendReqH() : nixlBackendReqH() {} - virtual ~nixlMooncakeBackendReqH(){} + + virtual ~nixlMooncakeBackendReqH() {} + uint64_t batch_id; size_t request_count; }; nixl_status_t -nixlMooncakeEngine::prepXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH *&handle, - const nixl_opt_b_args_t *opt_args) const { +nixlMooncakeEngine::prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { auto priv = new nixlMooncakeBackendReqH(); priv->batch_id = INVALID_BATCH; handle = priv; return NIXL_SUCCESS; } -nixl_status_t nixlMooncakeEngine::postXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args) const -{ - auto priv = (nixlMooncakeBackendReqH *) handle; +nixl_status_t +nixlMooncakeEngine::postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { + auto priv = (nixlMooncakeBackendReqH *)handle; int segment_id; { std::lock_guard lock(mutex_); const auto agent = connected_agents_.find(remote_agent); - if (agent == connected_agents_.end()) - return NIXL_ERR_INVALID_PARAM; + if (agent == connected_agents_.end()) return NIXL_ERR_INVALID_PARAM; segment_id = agent->second.segment_id; } if (local.descCount() != remote.descCount()) return NIXL_ERR_INVALID_PARAM; const static size_t kMaxRequestCount = 1024; if (priv->batch_id == INVALID_BATCH) { - uint64_t batch_id = allocateBatchID (engine_, kMaxRequestCount); + uint64_t batch_id = allocateBatchID(engine_, kMaxRequestCount); if (batch_id == INVALID_BATCH) { return NIXL_ERR_BACKEND; } @@ -256,21 +272,24 @@ nixl_status_t nixlMooncakeEngine::postXfer (const nixl_xfer_op_t &operation, request[index].length = local[index].len; request[index].target_id = segment_id; } - - // TODO: submitTransfer will fail when the total number of requests exceeded the - // batch size set for this batch ID. - int rc = submitTransfer(engine_, priv->batch_id, request, request_count); - delete []request; - if (rc) { - return NIXL_ERR_BACKEND; + int rc = 0; + if (opt_args->hasNotif) { + notify_msg_t notify_msg; + notify_msg.name = const_cast(local_agent_name_.c_str()); + notify_msg.msg = const_cast(opt_args->notifMsg.c_str()); + rc = submitTransferWithNotify(engine_, priv->batch_id, request, request_count, notify_msg); + } else { + rc = submitTransfer(engine_, priv->batch_id, request, request_count); } + delete[] request; + if (rc) return NIXL_ERR_BACKEND; priv->request_count += request_count; return NIXL_IN_PROG; } -nixl_status_t nixlMooncakeEngine::checkXfer (nixlBackendReqH* handle) const -{ - auto priv = (nixlMooncakeBackendReqH *) handle; +nixl_status_t +nixlMooncakeEngine::checkXfer(nixlBackendReqH *handle) const { + auto priv = (nixlMooncakeBackendReqH *)handle; bool has_failed = false; for (size_t index = 0; index < priv->request_count; ++index) { transfer_status_t status; @@ -284,18 +303,46 @@ nixl_status_t nixlMooncakeEngine::checkXfer (nixlBackendReqH* handle) const // Each batch_id has the batch size, and cannot process more requests // than the batch size. So, free the batch id here to workaround the issue // where the same nixlBackendReqH could be used to post multiple transfer. - freeBatchID (engine_, priv->batch_id); + freeBatchID(engine_, priv->batch_id); priv->batch_id = INVALID_BATCH; } return has_failed ? NIXL_ERR_BACKEND : NIXL_SUCCESS; } -nixl_status_t nixlMooncakeEngine::releaseReqH(nixlBackendReqH* handle) const -{ - auto priv = (nixlMooncakeBackendReqH *) handle; +nixl_status_t +nixlMooncakeEngine::releaseReqH(nixlBackendReqH *handle) const { + auto priv = (nixlMooncakeBackendReqH *)handle; if (priv->batch_id != INVALID_BATCH) { - freeBatchID (engine_, priv->batch_id); + freeBatchID(engine_, priv->batch_id); } delete priv; return NIXL_SUCCESS; } + +nixl_status_t +nixlMooncakeEngine::getNotifs(notif_list_t ¬if_list) { + if (notif_list.size() != 0) return NIXL_ERR_INVALID_PARAM; + int size = 0; + notify_msg_t *notify_msgs = getNotifsFromEngine(engine_, &size); + for (int i = 0; i < size; i++) { + notif_list.push_back(std::make_pair(notify_msgs[i].name, notify_msgs[i].msg)); + } + freeNotifsMsgBuf(notify_msgs, size); + return NIXL_SUCCESS; +} + +nixl_status_t +nixlMooncakeEngine::genNotif(const std::string &remote_agent, const std::string &msg) const { + int segment_id; + { + std::lock_guard lock(mutex_); + const auto agent = connected_agents_.find(remote_agent); + if (agent == connected_agents_.end()) return NIXL_ERR_INVALID_PARAM; + segment_id = agent->second.segment_id; + } + notify_msg_t notify_msg; + notify_msg.name = const_cast(local_agent_name_.c_str()); + notify_msg.msg = const_cast(msg.c_str()); + int ret = genNotifyInEngine(engine_, segment_id, notify_msg); + return nixl_status_t(ret); +} diff --git a/src/plugins/mooncake/mooncake_backend.h b/src/plugins/mooncake/mooncake_backend.h index bbd659a8e4..e76c22f550 100644 --- a/src/plugins/mooncake/mooncake_backend.h +++ b/src/plugins/mooncake/mooncake_backend.h @@ -27,78 +27,101 @@ #include "nixl.h" #include "backend/backend_engine.h" #include "common/str_tools.h" - #include "common/nixl_time.h" -#include "common/list_elem.h" #include "transfer_engine_c.h" class nixlMooncakeBackendMD; class nixlMooncakeEngine : public nixlBackendEngine { - public: - nixlMooncakeEngine(const nixlBackendInitParams* init_params); - ~nixlMooncakeEngine(); - - bool supportsRemote () const { return true; } - bool supportsLocal () const { return true; } - bool supportsNotif () const { return false; } - bool supportsProgTh () const { return false; } - - nixl_mem_list_t getSupportedMems () const; - - /* Object management */ - nixl_status_t getPublicData (const nixlBackendMD* meta, - std::string &str) const; - nixl_status_t getConnInfo(std::string &str) const; - nixl_status_t loadRemoteConnInfo (const std::string &remote_agent, - const std::string &remote_conn_info); - - nixl_status_t connect(const std::string &remote_agent); - nixl_status_t disconnect(const std::string &remote_agent); - - nixl_status_t registerMem (const nixlBlobDesc &mem, - const nixl_mem_t &nixl_mem, - nixlBackendMD* &out); - nixl_status_t deregisterMem (nixlBackendMD* meta); - - nixl_status_t loadLocalMD (nixlBackendMD* input, - nixlBackendMD* &output); - - nixl_status_t loadRemoteMD (const nixlBlobDesc &input, - const nixl_mem_t &nixl_mem, - const std::string &remote_agent, - nixlBackendMD* &output); - nixl_status_t unloadMD (nixlBackendMD* input); - - // Data transfer - nixl_status_t prepXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args=nullptr) const; - - nixl_status_t postXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args=nullptr) const; - - nixl_status_t checkXfer (nixlBackendReqH* handle) const; - nixl_status_t releaseReqH(nixlBackendReqH* handle) const; - - private: - struct AgentInfo { - int segment_id; - }; - - mutable std::mutex mutex_; - transfer_engine_t engine_; - std::string local_agent_name_; - std::unordered_map mem_reg_info_; - std::unordered_map connected_agents_; +public: + nixlMooncakeEngine(const nixlBackendInitParams *init_params); + ~nixlMooncakeEngine(); + + bool + supportsRemote() const { + return true; + } + + bool + supportsLocal() const { + return true; + } + + bool + supportsNotif() const { + return true; + } + + nixl_mem_list_t + getSupportedMems() const; + + /* Object management */ + nixl_status_t + getPublicData(const nixlBackendMD *meta, std::string &str) const; + nixl_status_t + getConnInfo(std::string &str) const; + nixl_status_t + loadRemoteConnInfo(const std::string &remote_agent, const std::string &remote_conn_info); + + nixl_status_t + connect(const std::string &remote_agent); + nixl_status_t + disconnect(const std::string &remote_agent); + + nixl_status_t + registerMem(const nixlBlobDesc &mem, const nixl_mem_t &nixl_mem, nixlBackendMD *&out); + nixl_status_t + deregisterMem(nixlBackendMD *meta); + + nixl_status_t + loadLocalMD(nixlBackendMD *input, nixlBackendMD *&output); + + nixl_status_t + loadRemoteMD(const nixlBlobDesc &input, + const nixl_mem_t &nixl_mem, + const std::string &remote_agent, + nixlBackendMD *&output); + nixl_status_t + unloadMD(nixlBackendMD *input); + + // Data transfer + nixl_status_t + prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const; + + nixl_status_t + postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const; + + nixl_status_t + checkXfer(nixlBackendReqH *handle) const; + nixl_status_t + releaseReqH(nixlBackendReqH *handle) const; + + nixl_status_t + getNotifs(notif_list_t ¬if_list); + nixl_status_t + genNotif(const std::string &remote_agent, const std::string &msg) const override; + +private: + struct AgentInfo { + int segment_id; + }; + + mutable std::mutex mutex_; + transfer_engine_t engine_; + std::string local_agent_name_; + std::unordered_map mem_reg_info_; + std::unordered_map connected_agents_; }; #endif diff --git a/src/plugins/mooncake/mooncake_plugin.cpp b/src/plugins/mooncake/mooncake_plugin.cpp index 30fe6e7436..9b4b85a49d 100644 --- a/src/plugins/mooncake/mooncake_plugin.cpp +++ b/src/plugins/mooncake/mooncake_plugin.cpp @@ -18,71 +18,31 @@ #include "backend/backend_plugin.h" #include "mooncake_backend.h" -// Plugin version information -static const char* PLUGIN_NAME = "Mooncake"; -static const char* PLUGIN_VERSION = "0.1.0"; - -// Function to create a new Mooncake backend engine instance -static nixlBackendEngine* create_mooncake_engine(const nixlBackendInitParams* init_params) { - return new nixlMooncakeEngine(init_params); -} - -static void destroy_mooncake_engine(nixlBackendEngine *engine) { - delete engine; -} - -// Function to get the plugin name -static const char* get_plugin_name() { - return PLUGIN_NAME; -} - -// Function to get the plugin version -static const char* get_plugin_version() { - return PLUGIN_VERSION; -} - -// Function to get backend options -static nixl_b_params_t get_backend_options() { +namespace { +nixl_b_params_t +get_mooncake_options() { nixl_b_params_t params; params["mooncake_devices"] = ""; return params; } +} // namespace -// Function to get supported backend mem types -static nixl_mem_list_t get_backend_mems() { - nixl_mem_list_t mems; - mems.push_back(DRAM_SEG); - mems.push_back(VRAM_SEG); - return mems; -} - -// Static plugin structure -static nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_mooncake_engine, - destroy_mooncake_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems -}; +// Plugin type alias for convenience +using mooncake_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_MOONCAKE - nixlBackendPlugin * -createStaticMooncakePlugin() { - return &plugin; // Return the static plugin instance +createStaticMOONCAKEPlugin() { + return mooncake_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "MOONCAKE", "0.1.0", get_mooncake_options(), {DRAM_SEG, VRAM_SEG}); } - #else - -// Plugin initialization function -extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - return &plugin; +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return mooncake_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "MOONCAKE", "0.1.0", get_mooncake_options(), {DRAM_SEG, VRAM_SEG}); } -// Plugin cleanup function -extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { - // Cleanup any resources if needed -} +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} #endif diff --git a/src/plugins/obj/README.md b/src/plugins/obj/README.md index 11eb297944..719f8acac2 100644 --- a/src/plugins/obj/README.md +++ b/src/plugins/obj/README.md @@ -47,6 +47,7 @@ Backend parameters are passed as a key-value map (`nixl_b_params_t`) when creati | `scheme` | HTTP scheme (`http` or `https`) | `https` | No | | `region` | AWS region for the S3 service | `us-east-1` | No | | `use_virtual_addressing` | Use virtual-hosted-style addressing (`true`/`false`) | `false` | No | +| `req_checksum` | Request checksum validation (`required`/`supported`) | - | No | \* If `access_key` and `secret_key` are not provided, the AWS SDK will attempt to use default credential providers (IAM roles, environment variables, credential files, etc.) diff --git a/src/plugins/obj/meson.build b/src/plugins/obj/meson.build index a1fd470006..69788f32f0 100644 --- a/src/plugins/obj/meson.build +++ b/src/plugins/obj/meson.build @@ -22,11 +22,14 @@ obj_sources = [ ] aws_s3 = dependency('aws-cpp-sdk-s3', static: false, required: false) -if aws_s3.found() +aws_core = dependency('aws-cpp-sdk-core', required: false, static: false) +if aws_s3.found() and aws_core.found() # By default aws-cpp-sdk sets c++11 compile flag for the whole project partial_aws_s3 = aws_s3.partial_dependency(compile_args: false, includes: true, link_args: true, links: true) - plugin_deps += [partial_aws_s3] + partial_aws_core = aws_core.partial_dependency(compile_args: false, includes: true, link_args: true, links: true) + plugin_deps += [partial_aws_s3, partial_aws_core] else + warning('AWS SDK dependencies not found, skipping OBJ plugin build') subdir_done() endif plugin_deps += [dependency('asio', required: true)] @@ -49,7 +52,8 @@ else include_directories: [nixl_inc_dirs, utils_inc_dirs], install: true, name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', 'echo "OBJ=' + obj_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', diff --git a/src/plugins/obj/obj_backend.h b/src/plugins/obj/obj_backend.h index 078cb5e082..c77c5bb543 100644 --- a/src/plugins/obj/obj_backend.h +++ b/src/plugins/obj/obj_backend.h @@ -46,11 +46,6 @@ class nixlObjEngine : public nixlBackendEngine { return false; } - bool - supportsProgTh() const override { - return false; - } - nixl_mem_list_t getSupportedMems() const override { return {OBJ_SEG, DRAM_SEG}; diff --git a/src/plugins/obj/obj_plugin.cpp b/src/plugins/obj/obj_plugin.cpp index 8e468b7892..df3883a321 100644 --- a/src/plugins/obj/obj_plugin.cpp +++ b/src/plugins/obj/obj_plugin.cpp @@ -20,72 +20,20 @@ #include "backend/backend_plugin.h" #include "common/nixl_log.h" -namespace { - -[[nodiscard]] nixlBackendEngine * -create_obj_engine(const nixlBackendInitParams *init_params) { - try { - return new nixlObjEngine(init_params); - } - catch (const std::exception &e) { - NIXL_ERROR << "Failed to create obj engine: " << e.what(); - return nullptr; - } -} - -void -destroy_obj_engine(nixlBackendEngine *engine) { - delete engine; -} - -[[nodiscard]] const char * -get_plugin_name() { - return "OBJ"; -} - -[[nodiscard]] const char * -get_plugin_version() { - return "0.1.0"; -} - -[[nodiscard]] nixl_b_params_t -get_backend_options() { - nixl_b_params_t params; - params["access_key"] = "AWS access key ID (required)"; - params["secret_key"] = "AWS secret access key (required)"; - params["session_token"] = "AWS session token (optional)"; - return params; -} - -[[nodiscard]] nixl_mem_list_t -get_backend_mems() { - return {DRAM_SEG, OBJ_SEG}; -} - -nixlBackendPlugin plugin = {NIXL_PLUGIN_API_VERSION, - create_obj_engine, - destroy_obj_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems}; -} // namespace +// Plugin type alias for convenience +using obj_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_OBJ - nixlBackendPlugin * -createStaticObjPlugin() { - return &plugin; // Return the static plugin instance +createStaticOBJPlugin() { + return obj_plugin_t::create(NIXL_PLUGIN_API_VERSION, "OBJ", "0.1.0", {}, {DRAM_SEG, OBJ_SEG}); } - -#else // !STATIC_PLUGIN_OBJ - +#else extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * nixl_plugin_init() { - return &plugin; + return obj_plugin_t::create(NIXL_PLUGIN_API_VERSION, "OBJ", "0.1.0", {}, {DRAM_SEG, OBJ_SEG}); } extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() {} - -#endif // !STATIC_PLUGIN_OBJ +#endif diff --git a/src/plugins/posix/meson.build b/src/plugins/posix/meson.build index bb9c20131a..2188f4e52d 100644 --- a/src/plugins/posix/meson.build +++ b/src/plugins/posix/meson.build @@ -70,7 +70,8 @@ else include_directories: [nixl_inc_dirs, utils_inc_dirs], install: true, name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', 'echo "POSIX=' + posix_backend_lib.full_path() + '" >> ' + plugin_build_dir + '/pluginlist', diff --git a/src/plugins/posix/posix_backend.cpp b/src/plugins/posix/posix_backend.cpp index 3031909a2c..93d3fd7c2c 100644 --- a/src/plugins/posix/posix_backend.cpp +++ b/src/plugins/posix/posix_backend.cpp @@ -126,9 +126,7 @@ nixlPosixBackendReqH::nixlPosixBackendReqH(const nixl_xfer_op_t &op, , queue_depth_(loc.descCount()) , queue_type_(getQueueType(params)) { if (queue_type_ == nixlPosixQueue::queue_t::UNSUPPORTED) { - throw exception( - absl::StrFormat("Unsupported backend type: %s", queue_type_), - NIXL_ERR_NOT_SUPPORTED); + throw exception(absl::StrFormat("Unsupported queue type"), NIXL_ERR_NOT_SUPPORTED); } if (local.descCount() == 0 || remote.descCount() == 0) { @@ -139,9 +137,8 @@ nixlPosixBackendReqH::nixlPosixBackendReqH(const nixl_xfer_op_t &op, nixl_status_t status = initQueues(); if (status != NIXL_SUCCESS) { - throw exception( - absl::StrFormat("Failed to initialize queues: %s", queue_type_), - status); + throw exception(absl::StrFormat("Failed to initialize queues: %s", to_string(queue_type_)), + status); } } @@ -156,7 +153,7 @@ nixl_status_t nixlPosixBackendReqH::initQueues() { queue = QueueFactory::createUringQueue(queue_depth_, operation); break; default: - NIXL_ERROR << absl::StrFormat("Invalid queue type: %s", queue_type_); + NIXL_ERROR << absl::StrFormat("Invalid queue type: %s", to_string(queue_type_)); return NIXL_ERR_INVALID_PARAM; } return NIXL_SUCCESS; @@ -206,11 +203,13 @@ nixlPosixEngine::nixlPosixEngine(const nixlBackendInitParams* init_params) , queue_type_(getQueueType(init_params->customParams)) { if (queue_type_ == nixlPosixQueue::queue_t::UNSUPPORTED) { initErr = true; - NIXL_ERROR << absl::StrFormat("Failed to initialize POSIX backend - requested backend not available: %s", - queue_type_); + NIXL_ERROR << absl::StrFormat( + "Failed to initialize POSIX backend - requested queue type not available: %s", + to_string(queue_type_)); return; } - NIXL_INFO << absl::StrFormat("POSIX backend initialized using %s backend", queue_type_); + NIXL_INFO << absl::StrFormat("POSIX backend initialized using queue type: %s", + to_string(queue_type_)); } nixl_status_t nixlPosixEngine::registerMem(const nixlBlobDesc &mem, @@ -248,7 +247,7 @@ nixl_status_t nixlPosixEngine::prepXfer(const nixl_xfer_op_t &operation, params["use_uring"] = "true"; break; default: - NIXL_ERROR << absl::StrFormat("Invalid queue type: %s", queue_type_); + NIXL_ERROR << absl::StrFormat("Invalid queue type: %s", to_string(queue_type_)); return NIXL_ERR_INVALID_PARAM; } diff --git a/src/plugins/posix/posix_backend.h b/src/plugins/posix/posix_backend.h index a2660f3d1b..b2777297b7 100644 --- a/src/plugins/posix/posix_backend.h +++ b/src/plugins/posix/posix_backend.h @@ -81,10 +81,6 @@ class nixlPosixEngine : public nixlBackendEngine { return false; } - bool supportsProgTh() const override { - return false; - } - nixl_mem_list_t getSupportedMems() const override { return {FILE_SEG, DRAM_SEG}; } diff --git a/src/plugins/posix/posix_plugin.cpp b/src/plugins/posix/posix_plugin.cpp index c1e70259aa..748f92da68 100644 --- a/src/plugins/posix/posix_plugin.cpp +++ b/src/plugins/posix/posix_plugin.cpp @@ -19,74 +19,22 @@ #include "posix_backend.h" #include "backend/backend_plugin.h" -// Function to create a new POSIX backend engine instance -static nixlBackendEngine* create_posix_engine(const nixlBackendInitParams* init_params) { - return new nixlPosixEngine(init_params); -} - -static void destroy_posix_engine(nixlBackendEngine *engine) { - delete engine; -} - -// Function to get the plugin name -static const char* get_plugin_name() { - return "POSIX"; -} - -// Function to get the plugin version -static const char* get_plugin_version() { - return "0.1.0"; -} - -// Function to get backend options -static nixl_b_params_t get_backend_options() { - nixl_b_params_t params; - return params; -} - -// Function to get supported backend mem types -static nixl_mem_list_t get_backend_mems() { - return {DRAM_SEG, FILE_SEG}; -} +// Plugin type alias for convenience +using posix_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_POSIX - -// Static plugin structure -static nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_posix_engine, - destroy_posix_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems -}; - -nixlBackendPlugin* createStaticPosixPlugin() { - return &plugin; // Return the static plugin instance +nixlBackendPlugin * +createStaticPOSIXPlugin() { + return posix_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "POSIX", "0.1.0", {}, {DRAM_SEG, FILE_SEG}); } - #else - -// Plugin initialization function -extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - try { - std::unique_ptr plugin = std::make_unique(); - plugin->create_engine = create_posix_engine; - plugin->destroy_engine = destroy_posix_engine; - plugin->get_plugin_name = get_plugin_name; - plugin->get_plugin_version = get_plugin_version; - plugin->get_backend_options = get_backend_options; - plugin->get_backend_mems = get_backend_mems; - plugin->api_version = NIXL_PLUGIN_API_VERSION; // Set the API version - return plugin.release(); - } catch (const std::exception& e) { - return nullptr; - } -} - -// Plugin cleanup function -extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return posix_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "POSIX", "0.1.0", {}, {DRAM_SEG, FILE_SEG}); } +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} #endif diff --git a/src/plugins/ucx/meson.build b/src/plugins/ucx/meson.build index 99e9daeb72..aebfa92f91 100644 --- a/src/plugins/ucx/meson.build +++ b/src/plugins/ucx/meson.build @@ -14,6 +14,7 @@ # limitations under the License. ucx_utils_dep = declare_dependency(link_with: ucx_utils_lib, include_directories: utils_inc_dirs ) +asio_dep = [dependency('asio', required: true)] compile_flags = [] if cuda_dep.found() @@ -23,7 +24,7 @@ endif if 'UCX' in static_plugins ucx_backend_lib = static_library('UCX', 'ucx_backend.cpp', 'ucx_backend.h', 'ucx_plugin.cpp', - dependencies: [nixl_infra, ucx_utils_dep, serdes_interface, cuda_dep, ucx_dep, thread_dep, nixl_common_dep], + dependencies: [nixl_infra, ucx_utils_dep, serdes_interface, cuda_dep, ucx_dep, thread_dep, nixl_common_dep, asio_dep], include_directories: nixl_inc_dirs, install: false, cpp_args : compile_flags, @@ -31,12 +32,13 @@ if 'UCX' in static_plugins else ucx_backend_lib = shared_library('UCX', 'ucx_backend.cpp', 'ucx_backend.h', 'ucx_plugin.cpp', - dependencies: [nixl_infra, ucx_utils_dep, serdes_interface, cuda_dep, ucx_dep, thread_dep, nixl_common_dep], + dependencies: [nixl_infra, ucx_utils_dep, serdes_interface, cuda_dep, ucx_dep, thread_dep, nixl_common_dep, asio_dep], include_directories: nixl_inc_dirs, install: true, cpp_args : compile_flags + ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + install_rpath: '$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', diff --git a/src/plugins/ucx/ucx_backend.cpp b/src/plugins/ucx/ucx_backend.cpp index 5cf5412a37..57b808f54a 100644 --- a/src/plugins/ucx/ucx_backend.cpp +++ b/src/plugins/ucx/ucx_backend.cpp @@ -19,12 +19,16 @@ #include "common/nixl_log.h" #include "serdes/serdes.h" #include "common/nixl_log.h" +#include "ucx/gpu_xfer_req_h.h" #include #include +#include #include #include #include "absl/strings/numbers.h" +#include "absl/strings/str_join.h" +#include #ifdef HAVE_CUDA @@ -233,8 +237,8 @@ void nixlUcxEngine::vramInitCtx() cudaCtx = std::make_unique(); } -int nixlUcxEngine::vramUpdateCtx(void *address, uint64_t devId, bool &restart_reqd) -{ +int +nixlUcxEngine::vramUpdateCtx(void *address, uint64_t dev_id, bool &restart_reqd) { int ret; bool was_updated; @@ -245,7 +249,7 @@ int nixlUcxEngine::vramUpdateCtx(void *address, uint64_t devId, bool &restart_r return 0; } - ret = cudaCtx->cudaUpdateCtxPtr(address, devId, was_updated); + ret = cudaCtx->cudaUpdateCtxPtr(address, dev_id, was_updated); if (ret) { return ret; } @@ -275,20 +279,42 @@ void nixlUcxEngine::vramFiniCtx() *****************************************/ -class nixlUcxIntReq : public nixlLinkElem { - private: - int _completed; - public: - std::unique_ptr amBuffer; +class nixlUcxIntReq { +public: + std::unique_ptr amBuffer; - nixlUcxIntReq() : nixlLinkElem() { - _completed = 0; - } + bool + is_complete() const { + return completed_; + } + + void + completed() { + completed_ = true; + } + + void + setConnection(ucx_connection_ptr_t conn) { + conn_ = conn; + } + + nixl_status_t + checkConnection(size_t ep_id) const { + NIXL_ASSERT(conn_) << "Connection is not set"; + return conn_->getEp(ep_id)->checkTxState(); + } - bool is_complete() const { return _completed; } - void completed() { _completed = 1; } +private: + bool completed_ = false; + ucx_connection_ptr_t conn_; }; +static void +nixlUcxReqSetConnection(nixlUcxReq req, ucx_connection_ptr_t conn) { + nixlUcxIntReq *req_int = reinterpret_cast(req); + req_int->setConnection(conn); +} + static void _internalRequestInit(void *request) { /* Initialize request in-place (aka "placement new")*/ @@ -314,8 +340,8 @@ static void _internalRequestReset(nixlUcxIntReq *req) { class nixlUcxBackendH : public nixlBackendReqH { private: - nixlUcxIntReq head; - const nixlUcxEngine ŋ + std::vector requests_; + nixlUcxWorker *worker; size_t worker_id; // Notification to be sent after completion of all requests @@ -332,92 +358,99 @@ class nixlUcxBackendH : public nixlBackendReqH { return notif; } - nixlUcxBackendH(const nixlUcxEngine &eng_, size_t worker_id_): eng(eng_), worker_id(worker_id_) {} + nixlUcxBackendH(nixlUcxWorker *worker, size_t worker_id) + : worker(worker), + worker_id(worker_id) {} - void append(nixlUcxIntReq *req) { - head.link(req); + void + reserve(size_t size) { + requests_.reserve(size); } - nixl_status_t release() - { - nixlUcxIntReq *req = head.next(); + void + append(nixlUcxIntReq *req) { + requests_.push_back(req); + } - if (!req) { - return NIXL_SUCCESS; - } + virtual bool + isComposite() const { + return false; + } - const auto &uw = eng.getWorker(worker_id); + virtual nixl_status_t + release() { // TODO: Error log: uncompleted requests found! Cancelling ... - while(req) { - nixlUcxIntReq *cur = req; - bool done = cur->is_complete(); - req = cur->unlink(); - if (!done) { + for (nixlUcxIntReq *req : requests_) { + if (!req->is_complete()) { // TODO: Need process this properly. // it may not be enough to cancel UCX request - uw->reqCancel((nixlUcxReq)cur); + worker->reqCancel((nixlUcxReq)req); } - _internalRequestReset(cur); - uw->reqRelease((nixlUcxReq)cur); + _internalRequestReset(req); + worker->reqRelease((nixlUcxReq)req); } + requests_.clear(); return NIXL_SUCCESS; } - - nixl_status_t status() - { - nixlUcxIntReq *req = head.next(); - nixl_status_t out_ret = NIXL_SUCCESS; - - if (NULL == req) { + virtual nixl_status_t + status() { + if (requests_.empty()) { /* No pending transmissions */ return NIXL_SUCCESS; } - const auto &uw = eng.getWorker(worker_id); - /* Maximum progress */ - while (uw->progress()); + while (worker->progress()) + ; /* Go over all request updating their status */ - while(req) { + nixl_status_t out_ret = NIXL_SUCCESS; + for (nixlUcxIntReq *req : requests_) { nixl_status_t ret; if (!req->is_complete()) { ret = ucx_status_to_nixl(ucp_request_check_status((nixlUcxReq)req)); switch (ret) { - case NIXL_SUCCESS: - /* Mark as completed */ - req->completed(); - break; - case NIXL_IN_PROG: - out_ret = NIXL_IN_PROG; - break; - default: - /* Any other ret value is ERR and will be returned */ - return ret; + case NIXL_SUCCESS: + /* Mark as completed */ + req->completed(); + break; + case NIXL_IN_PROG: + out_ret = NIXL_IN_PROG; + break; + default: + // Any other ret value is ERR and will be returned + nixl_status_t conn_status = req->checkConnection(worker_id); + return (conn_status == NIXL_SUCCESS) ? ret : conn_status; } } - req = req->next(); } - /* Remove completed requests keeping the first one as - request representative */ - req = head.unlink(); - while(req) { - nixlUcxIntReq *next_req = req->unlink(); + size_t incomplete_reqs = 0; + for (nixlUcxIntReq *req : requests_) { if (req->is_complete()) { _internalRequestReset(req); - uw->reqRelease((nixlUcxReq)req); + worker->reqRelease((nixlUcxReq)req); } else { - /* Enqueue back */ - append(req); + requests_[incomplete_reqs++] = req; } - req = next_req; } - + requests_.resize(incomplete_reqs); return out_ret; } + void + setWorker(nixlUcxWorker *worker, size_t worker_id) { + NIXL_ASSERT(this->worker == nullptr || worker == nullptr); + this->worker = worker; + this->worker_id = worker_id; + } + + nixlUcxWorker * + getWorker() const { + return worker; + } + size_t getWorkerId() const { return worker_id; } @@ -427,137 +460,674 @@ class nixlUcxBackendH : public nixlBackendReqH { * Progress thread management *****************************************/ -void nixlUcxEngine::progressFunc() -{ - using namespace nixlTime; +/* + * This class encapsulates a thread that polls one or multiple UCX workers + */ +class nixlUcxThread { +public: + nixlUcxThread(const nixlUcxEngine *engine, std::function init, size_t num_workers) + : engine_(engine), + init_(std::move(init)) { + workers_.reserve(num_workers); + } - vramApplyCtx(); + virtual ~nixlUcxThread() { + if (threadActive_) { + join(); + } + } - { - std::unique_lock lock(pthrActiveLock); - pthrActive = true; + void + start() { + NIXL_ASSERT(!threadActive_); + threadActive_ = std::make_unique>(); + auto active = threadActive_->get_future(); + thread_ = std::make_unique(std::ref(*this)); + active.wait(); } - pthrActiveCV.notify_one(); - // Set timeout event so that the main loop would progress all workers on first iteration - bool timeout = true; - bool pthrStop = false; - while (!pthrStop) { - for (size_t wid = 0; wid < pollFds.size() - 1; wid++) { - if (!(pollFds[wid].revents & POLLIN) && !timeout) - continue; - pollFds[wid].revents = 0; + virtual void + join() { + NIXL_ASSERT(threadActive_); + threadActive_.reset(); + thread_->join(); + } - bool made_progress = false; - nixl_status_t status; - const auto &uw = uws[wid]; - do { - while (uw->progress()) - made_progress = true; + virtual void + addWorker(nixlUcxWorker *worker, size_t worker_id) { + NIXL_ASSERT(workers_.size() < workers_.capacity()); + workers_.push_back(worker); + workerIds_.push_back(worker_id); + } - status = uw->arm(); - } while (status == NIXL_IN_PROG); - NIXL_ASSERT(status == NIXL_SUCCESS) << ", status: " << status; + const std::vector & + getWorkers() const { + return workers_; + } + + size_t + getWorkerId(size_t idx = 0) const { + return workerIds_[idx]; + } + + void + operator()() { + tlsThread() = this; + init_(); + threadActive_->set_value(); + run(); + } + + static nixlUcxThread *& + tlsThread() { + static thread_local nixlUcxThread *tls = nullptr; + return tls; + } + + static bool + isProgressThread(const nixlUcxEngine *engine) noexcept { + nixlUcxThread *thread = tlsThread(); + return thread && thread->engine_ == engine; + } + + friend std::ostream & + operator<<(std::ostream &os, const nixlUcxThread &thread) { + return os << "thread " << &thread << "{engine: " << thread.engine_ << ", worker_ids: [" + << absl::StrJoin(thread.workerIds_, ",") << "]}"; + } + +protected: + virtual void + run() = 0; + +private: + const nixlUcxEngine *engine_; + std::function init_; + std::vector workers_; + std::vector workerIds_; + std::unique_ptr thread_; + std::unique_ptr> threadActive_; +}; - if (made_progress && !wid) - notifProgress(); +class nixlUcxSharedThread : public nixlUcxThread { +public: + nixlUcxSharedThread(const nixlUcxEngine *engine, + std::function init, + size_t num_workers, + nixlTime::us_t delay) + : nixlUcxThread(engine, std::move(init), num_workers) { + if (pipe(controlPipe_) < 0) { + throw std::runtime_error("Couldn't create progress thread control pipe"); } - timeout = false; + // TODO: We need delay to manual periodic wakeup/polling as a temporary + // workaround for UCX bug (poll wouldn't wake up some fds in particular + // circumstances) + + // This will ensure that the resulting delay is at least 1ms and fits into int in order for + // it to be compatible with poll() + int delay_us = std::min((int)delay, std::numeric_limits::max()); + delay_ = std::chrono::ceil(std::chrono::microseconds(delay_us)); + + pollFds_.resize(num_workers + 1); + pollFds_.back() = {controlPipe_[0], POLLIN, 0}; + } + + ~nixlUcxSharedThread() { + close(controlPipe_[0]); + close(controlPipe_[1]); + } + + void + join() override { + const char signal = 'X'; + int ret = write(controlPipe_[1], &signal, sizeof(signal)); + if (ret < 0) NIXL_PERROR << "write to progress thread control pipe failed"; + nixlUcxThread::join(); + } + + void + addWorker(nixlUcxWorker *worker, size_t worker_id) override { + pollFds_[getWorkers().size()] = {worker->getEfd(), POLLIN, 0}; + nixlUcxThread::addWorker(worker, worker_id); + } - int ret; - while ((ret = poll(pollFds.data(), pollFds.size(), pthrDelay.count())) < 0) - NIXL_PTRACE << "Call to poll() was interrupted, retrying"; +protected: + void + run() override { + NIXL_DEBUG << "shared " << *this << " running"; + // Set timeout event so that the main loop would progress all workers on first iteration + bool timeout = true; + bool pthr_stop = false; + while (!pthr_stop) { + for (size_t i = 0; i < pollFds_.size() - 1; i++) { + if (!(pollFds_[i].revents & POLLIN) && !timeout) continue; + pollFds_[i].revents = 0; + nixlUcxWorker *worker = getWorkers()[i]; + do { + while (worker->progress()) + ; + } while (worker->arm() == NIXL_IN_PROG); + } + timeout = false; + + int ret; + while ((ret = poll(pollFds_.data(), pollFds_.size(), delay_.count())) < 0) + NIXL_PTRACE << "Call to poll() was interrupted, retrying"; - if (!ret) { - timeout = true; - } else if (pollFds.back().revents & POLLIN) { - pollFds.back().revents = 0; + if (!ret) { + timeout = true; + } else if (pollFds_.back().revents & POLLIN) { + pollFds_.back().revents = 0; - char signal; - int ret = read(pollFds.back().fd, &signal, sizeof(signal)); - if (ret < 0) - NIXL_PERROR << "read() on control pipe failed"; + char signal; + int ret = read(pollFds_.back().fd, &signal, sizeof(signal)); + if (ret < 0) NIXL_PERROR << "read() on control pipe failed"; - pthrStop = true; + pthr_stop = true; + } } + + NIXL_DEBUG << "shared " << *this << " exiting"; } -} -void nixlUcxEngine::progressThreadStart() -{ - { - std::unique_lock lock(pthrActiveLock); - pthrActive = false; +private: + std::chrono::milliseconds delay_; + int controlPipe_[2]; + std::vector pollFds_; +}; + +nixlUcxThreadEngine::nixlUcxThreadEngine(const nixlBackendInitParams &init_params) + : nixlUcxEngine(init_params) { + if (!nixlUcxMtLevelIsSupported(nixl_ucx_mt_t::WORKER)) { + throw std::invalid_argument("UCX library does not support multi-threading"); } - if (!pthrOn) { - // not enabled - return; + size_t num_workers = getWorkers().size(); + thread_ = std::make_unique( + this, [this]() { nixlUcxEngine::vramApplyCtx(); }, num_workers, init_params.pthrDelay); + for (size_t i = 0; i < num_workers; i++) { + thread_->addWorker(getWorkers()[i].get(), i); } + thread_->start(); +} - pthr = std::thread(&nixlUcxEngine::progressFunc, this); +nixlUcxThreadEngine::~nixlUcxThreadEngine() { + thread_->join(); +} - std::unique_lock lock(pthrActiveLock); - pthrActiveCV.wait(lock, [&]{ return pthrActive; }); +int +nixlUcxThreadEngine::vramApplyCtx() { + thread_->join(); + thread_->start(); + return nixlUcxEngine::vramApplyCtx(); } -void nixlUcxEngine::progressThreadStop() -{ - if (!pthrOn) { - // not enabled - return; +void +nixlUcxThreadEngine::appendNotif(std::string remote_name, std::string msg) { + if (nixlUcxThread::isProgressThread(this)) { + /* Append to the private list to allow batching */ + const std::lock_guard lock(notifMtx_); + notifPthr_.push_back(std::make_pair(std::move(remote_name), std::move(msg))); + } else { + nixlUcxEngine::appendNotif(std::move(remote_name), std::move(msg)); } - - const char signal = 'X'; - int ret = write(pthrControlPipe[1], &signal, sizeof(signal)); - if (ret < 0) - NIXL_PERROR << "write to progress thread control pipe failed"; - pthr.join(); } -void nixlUcxEngine::progressThreadRestart() -{ - progressThreadStop(); - progressThreadStart(); +nixl_status_t +nixlUcxThreadEngine::getNotifs(notif_list_t ¬if_list) { + if (!notif_list.empty()) return NIXL_ERR_INVALID_PARAM; + + getNotifsImpl(notif_list); + const std::lock_guard lock(notifMtx_); + moveNotifList(notifPthr_, notif_list); + return NIXL_SUCCESS; } /**************************************** - * Constructor/Destructor -*****************************************/ + * Threadpool engine + ****************************************/ -nixlUcxEngine::nixlUcxEngine(const nixlBackendInitParams *init_params) - : nixlBackendEngine(init_params), - pthrControlPipe{0, 0} { - size_t numWorkers; - std::vector devs; /* Empty vector */ - nixl_b_params_t* custom_params = init_params->customParams; +struct nixlUcxBackendSharedState; + +/* + * This class represents a chunk of a composite request. + * It is used to encapsulate a batch of requests (subset of the larger batch) + * performed by a dedicated worker thread of threadpool. It holds a shared state + * with the main request to track its completion status and control the lifetime. + */ +class nixlUcxChunkBackendH : public nixlUcxBackendH { +public: + nixlUcxChunkBackendH() : nixlUcxBackendH(nullptr, UINT64_MAX) {} + + void + startXfer(const std::shared_ptr &shared_state, + nixlUcxWorker *worker, + size_t worker_id) { + NIXL_ASSERT(sharedState_.get() == nullptr); + sharedState_ = shared_state; + setWorker(worker, worker_id); + } - if (init_params->enableProgTh) { - if (!nixlUcxMtLevelIsSupported(nixl_ucx_mt_t::WORKER)) { - throw std::invalid_argument("UCX library does not support multi-threading"); + void + complete(nixl_status_t status); + + nixl_status_t + status() override; + + friend std::ostream & + operator<<(std::ostream &os, const nixlUcxChunkBackendH &chunk) { + return os << "chunk " << &chunk << "{worker_id: " << chunk.getWorkerId() + << ", state: " << chunk.sharedState_.get() << "}"; + } + +private: + std::shared_ptr sharedState_; +}; + +/* + * This class represents a shared state between a main request and all of its + * chunks. It is used to track the completion status of the request and the + * number of pending requests, and to control the lifetime of the chunks. + */ +struct nixlUcxBackendSharedState { + std::atomic status; + std::atomic pendingReqs; + std::vector chunks; + + nixlUcxBackendSharedState() : status(NIXL_SUCCESS), pendingReqs(0) {} + + friend std::ostream & + operator<<(std::ostream &os, const nixlUcxBackendSharedState &state) { + return os << "state " << &state << "{status: " << state.status.load() + << ", pending=" << state.pendingReqs.load() << "}"; + } +}; + +void +nixlUcxChunkBackendH::complete(nixl_status_t status) { + NIXL_ASSERT(sharedState_.get() != nullptr); + if (status != NIXL_SUCCESS) { + nixlUcxBackendH::release(); + sharedState_->status.store(status); + } + sharedState_->pendingReqs.fetch_sub(1); + NIXL_TRACE << *this << " completed with status: " << status << ", " << *sharedState_; + setWorker(nullptr, UINT64_MAX); + sharedState_.reset(); +} + +nixl_status_t +nixlUcxChunkBackendH::status() { + // First check if entire request was cancelled or failed + nixl_status_t status = sharedState_->status.load(); + if (status == NIXL_SUCCESS) { + status = nixlUcxBackendH::status(); + } + return status; +} + +/* + * This class represents a composite request handle for a UCX backend. + * It is used to encapsulate multiple parallel requests performed by dedicated + * worker threads of threadpool, with a single request handle, that it returned + * to the user. + */ +class nixlUcxCompositeBackendH : public nixlUcxBackendH { +public: + nixlUcxCompositeBackendH(nixlUcxWorker *worker, + size_t worker_id, + size_t chunk_size, + size_t num_chunks) + : nixlUcxBackendH(worker, worker_id), + sharedState_(std::make_shared()), + chunkSize_(chunk_size) { + sharedState_->chunks.resize(num_chunks); + } + + size_t + getChunkSize() const { + return chunkSize_; + } + + size_t + getNumChunks() const { + return sharedState_ ? sharedState_->chunks.size() : 0; + } + + void + startXfer() { + NIXL_ASSERT(sharedState_->pendingReqs.load() == 0); + sharedState_->status.store(NIXL_SUCCESS); + sharedState_->pendingReqs.store(getNumChunks()); + } + + nixlUcxChunkBackendH * + startChunk(size_t idx, nixlUcxWorker *worker, size_t worker_id) { + nixlUcxChunkBackendH *chunk = &sharedState_->chunks[idx]; + chunk->startXfer(sharedState_, worker, worker_id); + return chunk; + } + + bool + isComposite() const override { + return true; + } + + nixl_status_t + release() override { + NIXL_TRACE << *this << " releasing"; + nixl_status_t status = nixlUcxBackendH::release(); + if (sharedState_) { + // Set failed status to stop progress chunks + sharedState_->status.store(NIXL_ERR_NOT_FOUND); + // Reset shared state - it will be effectively released when the last chunk + // resets the shared state pointer + sharedState_.reset(); } - if (pipe(pthrControlPipe) < 0) { - throw std::runtime_error("Couldn't create progress thread control pipe"); + return status; + } + + nixl_status_t + status() override { + while (getWorker()->progress()) + ; + + if (sharedState_->pendingReqs.load()) { + return NIXL_IN_PROG; } - // This will ensure that the resulting delay is at least 1ms and fits into int in order for - // it to be compatible with poll() - pthrDelay = std::chrono::ceil( - std::chrono::microseconds(init_params->pthrDelay < std::numeric_limits::max() ? - init_params->pthrDelay : - std::numeric_limits::max())); - pthrOn = true; + nixl_status_t status = nixlUcxBackendH::status(); + if (status != NIXL_SUCCESS) { + return status; + } + + return sharedState_->status.load(); + } + + friend std::ostream & + operator<<(std::ostream &os, const nixlUcxCompositeBackendH &handle) { + os << "composite handle " << &handle << "{chunks: " << handle.getNumChunks(); + if (handle.sharedState_) { + os << ", " << *handle.sharedState_; + } else { + os << ", state: nullptr"; + } + return os << "}}"; + } + +private: + std::shared_ptr sharedState_; + size_t chunkSize_; +}; + +class nixlUcxDedicatedThread : public nixlUcxThread { +public: + nixlUcxDedicatedThread(nixlUcxEngine *engine, std::function init, asio::io_context &io) + : nixlUcxThread(engine, std::move(init), 1), + io_(io) {} + + static nixlUcxDedicatedThread * + getDedicatedThread() { + return (nixlUcxDedicatedThread *)tlsThread(); + } + + void + addRequest(nixlUcxChunkBackendH *handle) { + requests_.push_back(handle); + } + +protected: + void + run() override { + auto guard = asio::make_work_guard(io_); + NIXL_DEBUG << "dedicated " << *this << " running"; + + while (!io_.stopped()) { + if (!requests_.empty()) { + io_.poll_one(); + } else { + NIXL_TRACE << "dedicated " << *this << " waiting for requests"; + io_.run_one(); + } + + if (requests_.empty()) { + continue; + } + + for (auto it = requests_.begin(); it != requests_.end();) { + nixl_status_t status = (*it)->status(); + if (status != NIXL_IN_PROG) { + NIXL_TRACE << "dedicated " << *this << " completing " << *(*it) + << " with status: " << status; + (*it)->complete(status); + it = requests_.erase(it); + } else { + ++it; + } + } + } + + if (!requests_.empty()) { + NIXL_WARN << "dedicated " << *this << " dropping " << requests_.size() + << " requests on exit"; + for (auto it = requests_.begin(); it != requests_.end();) { + NIXL_INFO << "dropping " << *(*it); + (*it)->complete(NIXL_ERR_BACKEND); + } + requests_.clear(); + } + + NIXL_DEBUG << "dedicated " << *this << " exiting"; + } + +private: + asio::io_context &io_; + std::vector requests_; +}; + +nixlUcxThreadPoolEngine::nixlUcxThreadPoolEngine(const nixlBackendInitParams &init_params) + : nixlUcxEngine(init_params) { + size_t num_threads = nixl_b_params_get(init_params.customParams, "num_threads", 0); + numSharedWorkers_ = getWorkers().size() - num_threads; + NIXL_ASSERT(numSharedWorkers_ > 0); + + splitBatchSize_ = nixl_b_params_get(init_params.customParams, "split_batch_size", 1024); + + auto init = [this]() { nixlUcxEngine::vramApplyCtx(); }; + + if (init_params.enableProgTh) { + sharedThread_ = std::make_unique( + this, init, numSharedWorkers_, init_params.pthrDelay); + for (size_t i = 0; i < numSharedWorkers_; i++) { + sharedThread_->addWorker(getWorkers()[i].get(), i); + } + sharedThread_->start(); + } + + if (num_threads > 0) { + io_.reset(new asio::io_context()); + dedicatedThreads_.reserve(num_threads); + for (size_t i = 0; i < num_threads; ++i) { + size_t worker_id = numSharedWorkers_ + i; + dedicatedThreads_.emplace_back( + std::make_unique(this, init, *io_)); + dedicatedThreads_.back()->addWorker(getWorker(worker_id).get(), worker_id); + dedicatedThreads_.back()->start(); + } + } +} + +nixlUcxThreadPoolEngine::~nixlUcxThreadPoolEngine() { + if (sharedThread_) { + sharedThread_->join(); + } + + if (io_) { + io_->stop(); + for (auto &thread : dedicatedThreads_) { + thread->join(); + } + } +} + +nixl_status_t +nixlUcxThreadPoolEngine::prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { + size_t batch_size = local.descCount(); + if (batch_size < splitBatchSize_) { + return nixlUcxEngine::prepXfer(operation, local, remote, remote_agent, handle, opt_args); + } + + size_t chunk_size = std::max(batch_size / dedicatedThreads_.size(), splitBatchSize_); + size_t num_chunks = (batch_size + chunk_size - 1) / chunk_size; + + size_t worker_id = getWorkerId(); + auto comp_handle = + new nixlUcxCompositeBackendH(getWorker(worker_id).get(), worker_id, chunk_size, num_chunks); + NIXL_TRACE << "created " << *comp_handle; + handle = comp_handle; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlUcxThreadPoolEngine::sendXferRange(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *handle, + size_t start_idx, + size_t end_idx) const { + nixlUcxBackendH *int_handle = static_cast(handle); + if (!int_handle->isComposite()) { + return nixlUcxEngine::sendXferRange( + operation, local, remote, remote_agent, handle, start_idx, end_idx); + } + + nixlUcxCompositeBackendH *comp_handle = static_cast(int_handle); + comp_handle->startXfer(); + size_t chunk_size = comp_handle->getChunkSize(); + NIXL_TRACE << "sending " << *comp_handle; + + std::promise promise; + std::future future = promise.get_future(); + std::atomic remaining{comp_handle->getNumChunks()}; + std::atomic status{NIXL_SUCCESS}; + + for (size_t i = 0; i < comp_handle->getNumChunks(); i++) { + io_->post([&, i]() { + auto thread = nixlUcxDedicatedThread::getDedicatedThread(); + NIXL_ASSERT(thread != nullptr); + + nixlUcxChunkBackendH *chunk_handle = + comp_handle->startChunk(i, thread->getWorkers()[0], thread->getWorkerId()); + NIXL_TRACE << "dedicated " << *thread << " starting " << *chunk_handle; + + size_t start_idx = i * chunk_size; + size_t end_idx = std::min(start_idx + chunk_size, (size_t)local.descCount()); + nixl_status_t ret = nixlUcxEngine::sendXferRange( + operation, local, remote, remote_agent, chunk_handle, start_idx, end_idx); + if (ret != NIXL_SUCCESS) { + status.store(ret); + chunk_handle->complete(ret); + } else { + NIXL_TRACE << "dedicated " << *thread << " sent " << *chunk_handle; + thread->addRequest(chunk_handle); + } + + if (remaining.fetch_sub(1) == 1) { + promise.set_value(); + } + }); + } + + future.wait(); + NIXL_TRACE << "sent " << *comp_handle << " with status: " << status.load(); + return status.load(); +} + +int +nixlUcxThreadPoolEngine::vramApplyCtx() { + // TODO: Check if UCX can handle context change at runtime + if (sharedThread_) { + sharedThread_->join(); + sharedThread_->start(); + } + if (io_) { + io_->stop(); + for (auto &thread : dedicatedThreads_) { + thread->join(); + } + io_->restart(); + for (auto &thread : dedicatedThreads_) { + thread->start(); + } + } + return nixlUcxEngine::vramApplyCtx(); +} + +void +nixlUcxThreadPoolEngine::appendNotif(std::string remote_name, std::string msg) { + if (nixlUcxThread::isProgressThread(this)) { + std::lock_guard lock(notifMutex_); + notifThread_.emplace_back(std::move(remote_name), std::move(msg)); } else { - pthrOn = false; + nixlUcxEngine::appendNotif(std::move(remote_name), std::move(msg)); } +} + +nixl_status_t +nixlUcxThreadPoolEngine::getNotifs(notif_list_t ¬if_list) { + if (!notif_list.empty()) return NIXL_ERR_INVALID_PARAM; + + if (!sharedThread_) { + progress(); + } + + getNotifsImpl(notif_list); + std::lock_guard lock(notifMutex_); + moveNotifList(notifThread_, notif_list); + return NIXL_SUCCESS; +} + +/**************************************** + * Constructor/Destructor + *****************************************/ + +std::unique_ptr +nixlUcxEngine::create(const nixlBackendInitParams &init_params) { + nixlUcxEngine *engine; + size_t num_threads = nixl_b_params_get(init_params.customParams, "num_threads", 0); + if (num_threads > 0) { + engine = new nixlUcxThreadPoolEngine(init_params); + } else if (init_params.enableProgTh) { + engine = new nixlUcxThreadEngine(init_params); + } else { + engine = new nixlUcxEngine(init_params); + } + return std::unique_ptr(engine); +} + +nixlUcxEngine::nixlUcxEngine(const nixlBackendInitParams &init_params) + : nixlBackendEngine(&init_params), + sharedWorkerIndex_(1) { + std::vector devs; /* Empty vector */ + nixl_b_params_t *custom_params = init_params.customParams; if (custom_params->count("device_list")!=0) devs = str_split((*custom_params)["device_list"], ", "); - const auto num_workers_iter = custom_params->find("num_workers"); - if (num_workers_iter == custom_params->end() || !absl::SimpleAtoi(num_workers_iter->second, &numWorkers)) - numWorkers = 1; + size_t num_workers = nixl_b_params_get(custom_params, "num_workers", 1); + size_t num_threads = nixl_b_params_get(custom_params, "num_threads", 0); + + if (num_workers <= num_threads) { + /* There must be at least one shared worker */ + num_workers = num_threads + 1; + } ucp_err_handling_mode_t err_handling_mode; const auto err_handling_mode_it = @@ -572,28 +1142,16 @@ nixlUcxEngine::nixlUcxEngine(const nixlBackendInitParams *init_params) sizeof(nixlUcxIntReq), _internalRequestInit, _internalRequestFini, - pthrOn, - numWorkers, - init_params->syncMode); + init_params.enableProgTh, + num_workers, + init_params.syncMode); - for (size_t i = 0; i < numWorkers; i++) { + for (size_t i = 0; i < num_workers; i++) { uws.emplace_back(std::make_unique(*uc, err_handling_mode)); } - workerAddr = uws.front()->epAddr(); - - if (pthrOn) { - for (auto &uw: uws) { - pollFds.push_back({uw->getEfd(), POLLIN, 0}); - } - pollFds.push_back({pthrControlPipe[0], POLLIN, 0}); - } - - // TODO: in case of UCX error handling is enabled, we can clean up AM based connections error - // handling, if user requested disabled error handling, we dont care about it. auto &uw = uws.front(); - uw->regAmCallback(CONN_CHECK, connectionCheckAmCb, this); - uw->regAmCallback(DISCONNECT, connectionTermAmCb, this); + workerAddr = uw->epAddr(); uw->regAmCallback(NOTIF_STR, notifAmCb, this); // Temp fixup @@ -606,7 +1164,6 @@ nixlUcxEngine::nixlUcxEngine(const nixlBackendInitParams *init_params) m_cudaPrimaryCtx = std::make_shared(); vramInitCtx(); - progressThreadStart(); } nixl_mem_list_t nixlUcxEngine::getSupportedMems () const { @@ -616,18 +1173,16 @@ nixl_mem_list_t nixlUcxEngine::getSupportedMems () const { return mems; } +static std::unordered_map & +tlsSharedWorkerMap() { + static thread_local std::unordered_map map; + return map; +} + // Through parent destructor the unregister will be called. nixlUcxEngine::~nixlUcxEngine() { - progressThreadStop(); - if (pthrOn) { - for (const auto pthr_control_pipe : pthrControlPipe) { - if (pthr_control_pipe != 0) { - close(pthr_control_pipe); - } - } - } - vramFiniCtx(); + tlsSharedWorkerMap().erase(this); } /**************************************** @@ -638,124 +1193,29 @@ nixl_status_t nixlUcxEngine::checkConn(const std::string &remote_agent) { return remoteConnMap.count(remote_agent) ? NIXL_SUCCESS : NIXL_ERR_NOT_FOUND; } -nixl_status_t nixlUcxEngine::endConn(const std::string &remote_agent) { - - auto search = remoteConnMap.find(remote_agent); - - if(search == remoteConnMap.end()) { - return NIXL_ERR_NOT_FOUND; - } - - //thread safety? - remoteConnMap.erase(search); - - return NIXL_SUCCESS; -} - nixl_status_t nixlUcxEngine::getConnInfo(std::string &str) const { str = workerAddr; return NIXL_SUCCESS; } -ucs_status_t -nixlUcxEngine::connectionCheckAmCb(void *arg, const void *header, - size_t header_length, void *data, - size_t length, - const ucp_am_recv_param_t *param) -{ - std::string remote_agent( (char*) data, length); - nixlUcxEngine* engine = (nixlUcxEngine*) arg; - - NIXL_ASSERT(!(param->recv_attr & UCP_AM_RECV_ATTR_FLAG_RNDV)); - NIXL_ASSERT(header_length == 0) << "header_length " << header_length; - - if(engine->checkConn(remote_agent)) { - NIXL_ERROR << "Received connect AM from agent we don't recognize: " << remote_agent; - return UCS_OK; - } - - return UCS_OK; -} - -ucs_status_t -nixlUcxEngine::connectionTermAmCb (void *arg, const void *header, - size_t header_length, void *data, - size_t length, - const ucp_am_recv_param_t *param) -{ - std::string remote_agent( (char*) data, length); - - NIXL_ASSERT(!(param->recv_attr & UCP_AM_RECV_ATTR_FLAG_RNDV)); - NIXL_ASSERT(header_length == 0) << "header_length " << header_length; - -/* - // TODO: research UCX connection logic and fix. - nixlUcxEngine* engine = (nixlUcxEngine*) arg; - if(NIXL_SUCCESS != engine->endConn(remote_agent)) { - //TODO: received connect AM from agent we don't recognize - return UCS_ERR_INVALID_PARAM; - } -*/ - return UCS_OK; -} - nixl_status_t nixlUcxEngine::connect(const std::string &remote_agent) { if(remote_agent == localAgent) { return loadRemoteConnInfo(remote_agent, workerAddr); } - const auto search = remoteConnMap.find(remote_agent); - - if(search == remoteConnMap.end()) { - return NIXL_ERR_NOT_FOUND; - } - - bool error = false; - nixl_status_t ret = NIXL_SUCCESS; - std::vector reqs; - for (size_t i = 0; i < uws.size(); i++) { - reqs.emplace_back(); - ret = search->second->getEp(i)->sendAm(CONN_CHECK, NULL, 0, - (void*) localAgent.data(), localAgent.size(), - UCP_AM_SEND_FLAG_EAGER, reqs.back()); - if(ret < 0) { - error = true; - break; - } - } - - //wait for AM to send - ret = NIXL_IN_PROG; - for (size_t i = 0; i < reqs.size(); i++) - while(ret == NIXL_IN_PROG) - ret = getWorker(i)->test(reqs[i]); - return error ? NIXL_ERR_BACKEND : NIXL_SUCCESS; + return (remoteConnMap.find(remote_agent) == remoteConnMap.end()) ? NIXL_ERR_NOT_FOUND : + NIXL_SUCCESS; } nixl_status_t nixlUcxEngine::disconnect(const std::string &remote_agent) { - if (remote_agent != localAgent) { - auto search = remoteConnMap.find(remote_agent); - - if(search == remoteConnMap.end()) { - return NIXL_ERR_NOT_FOUND; - } + auto search = remoteConnMap.find(remote_agent); - nixl_status_t ret = NIXL_SUCCESS; - for (size_t i = 0; i < uws.size(); i++) { - if (search->second->getEp(i)->checkTxState() == NIXL_SUCCESS) { - nixlUcxReq req; - ret = search->second->getEp(i)->sendAm(DISCONNECT, NULL, 0, - (void*) localAgent.data(), localAgent.size(), - UCP_AM_SEND_FLAG_EAGER, req); - //don't care - if (ret == NIXL_IN_PROG) - getWorker(i)->reqRelease(req); - } - } + if (search == remoteConnMap.end()) { + return NIXL_ERR_NOT_FOUND; } - endConn(remote_agent); - + // thread safety? + remoteConnMap.erase(search); return NIXL_SUCCESS; } @@ -807,8 +1267,6 @@ nixl_status_t nixlUcxEngine::registerMem (const nixlBlobDesc &mem, //TODO Add to logging } if (need_restart) { - progressThreadRestart(); - // set the ctx for main thread vramApplyCtx(); } } @@ -908,23 +1366,39 @@ nixl_status_t nixlUcxEngine::unloadMD (nixlBackendMD* input) { * Data movement *****************************************/ -static nixl_status_t _retHelper(nixl_status_t ret, nixlUcxBackendH *hndl, nixlUcxReq &req) -{ +static nixl_status_t +_retHelper(nixl_status_t ret, nixlUcxBackendH *hndl, nixlUcxReq &req, ucx_connection_ptr_t conn) { /* if transfer wasn't immediately completed */ switch(ret) { - case NIXL_IN_PROG: - hndl->append((nixlUcxIntReq*)req); - case NIXL_SUCCESS: - // Nothing to do - break; - default: - // Error. Release all previously initiated ops and exit: - hndl->release(); - return NIXL_ERR_BACKEND; + case NIXL_IN_PROG: + // TODO: this cast does not look safe + // We need to allocate a vector of nixlUcxIntReq and set nixlUcxReqt + hndl->append((nixlUcxIntReq *)req); + nixlUcxReqSetConnection(req, conn); + case NIXL_SUCCESS: + // Nothing to do + break; + default: + // Error. Release all previously initiated ops and exit: + hndl->release(); + return ret; } + return NIXL_SUCCESS; } +size_t +nixlUcxEngine::getWorkerId() const { + auto it = tlsSharedWorkerMap().find(this); + if (it == tlsSharedWorkerMap().end()) { + size_t index = sharedWorkerIndex_.fetch_add(1) % getSharedWorkersSize(); + it = tlsSharedWorkerMap().emplace(this, index).first; + NIXL_DEBUG << "engine " << this << " bound shared worker " << index << " to thread " + << std::this_thread::get_id(); + } + return it->second; +} + nixl_status_t nixlUcxEngine::prepXfer (const nixl_xfer_op_t &operation, const nixl_meta_dlist_t &local, const nixl_meta_dlist_t &remote, @@ -932,10 +1406,17 @@ nixl_status_t nixlUcxEngine::prepXfer (const nixl_xfer_op_t &operation, nixlBackendReqH* &handle, const nixl_opt_b_args_t* opt_args) const { + if (local.descCount() == 0 || remote.descCount() == 0) { + NIXL_ERROR << "Local or remote descriptor list is empty"; + return NIXL_ERR_INVALID_PARAM; + } + /* TODO: try to get from a pool first */ - nixlUcxBackendH *intHandle = new nixlUcxBackendH(*this, getWorkerId()); + size_t worker_id = getWorkerId(); + auto *ucx_handle = new nixlUcxBackendH(getWorker(worker_id).get(), worker_id); + + handle = ucx_handle; - handle = (nixlBackendReqH*)intHandle; return NIXL_SUCCESS; } @@ -995,35 +1476,33 @@ nixl_status_t nixlUcxEngine::estimateXferCost (const nixl_xfer_op_t &operation, return NIXL_SUCCESS; } -nixl_status_t nixlUcxEngine::postXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args) const -{ - size_t lcnt = local.descCount(); - size_t rcnt = remote.descCount(); - size_t i; - nixl_status_t ret; +nixl_status_t +nixlUcxEngine::sendXferRange(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *handle, + size_t start_idx, + size_t end_idx) const { nixlUcxBackendH *intHandle = (nixlUcxBackendH *)handle; nixlUcxPrivateMetadata *lmd; nixlUcxPublicMetadata *rmd; + nixl_status_t ret; nixlUcxReq req; size_t workerId = intHandle->getWorkerId(); - if (lcnt != rcnt) { - return NIXL_ERR_INVALID_PARAM; - } + // Reserve space for the requests, +2 for flush and completion + intHandle->reserve(end_idx - start_idx + 2); - for(i = 0; i < lcnt; i++) { + for (size_t i = start_idx; i < end_idx; i++) { void *laddr = (void*) local[i].addr; size_t lsize = local[i].len; - void *raddr = (void*) remote[i].addr; + uint64_t raddr = (uint64_t)remote[i].addr; size_t rsize = remote[i].len; lmd = (nixlUcxPrivateMetadata*) local[i].metadataP; rmd = (nixlUcxPublicMetadata*) remote[i].metadataP; + auto &ep = rmd->conn->getEp(workerId); if (lsize != rsize) { return NIXL_ERR_INVALID_PARAM; @@ -1031,16 +1510,16 @@ nixl_status_t nixlUcxEngine::postXfer (const nixl_xfer_op_t &operation, switch (operation) { case NIXL_READ: - ret = rmd->conn->getEp(workerId)->read((uint64_t) raddr, rmd->getRkey(workerId), laddr, lmd->mem, lsize, req); + ret = ep->read(raddr, rmd->getRkey(workerId), laddr, lmd->mem, lsize, req); break; case NIXL_WRITE: - ret = rmd->conn->getEp(workerId)->write(laddr, lmd->mem, (uint64_t) raddr, rmd->getRkey(workerId), lsize, req); + ret = ep->write(laddr, lmd->mem, raddr, rmd->getRkey(workerId), lsize, req); break; default: return NIXL_ERR_INVALID_PARAM; } - if (_retHelper(ret, intHandle, req)) { + if (_retHelper(ret, intHandle, req, rmd->conn)) { return ret; } } @@ -1049,23 +1528,55 @@ nixl_status_t nixlUcxEngine::postXfer (const nixl_xfer_op_t &operation, * Flush keeps intHandle non-empty until the operation is actually * completed, which can happen after local requests completion. */ - rmd = (nixlUcxPublicMetadata*) remote[0].metadataP; + rmd = (nixlUcxPublicMetadata *)remote[0].metadataP; ret = rmd->conn->getEp(workerId)->flushEp(req); - if (_retHelper(ret, intHandle, req)) { + + if (_retHelper(ret, intHandle, req, rmd->conn)) { + return ret; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlUcxEngine::postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args) const { + size_t lcnt = local.descCount(); + size_t rcnt = remote.descCount(); + nixlUcxBackendH *int_handle = static_cast(handle); + nixl_status_t ret; + + if (lcnt != rcnt) { + NIXL_ERROR << "Local (" << lcnt << ") and remote (" << rcnt + << ") descriptor lists differ in size"; + return NIXL_ERR_INVALID_PARAM; + } + + // TODO: assert that handle is empty/completed, as we can't post request before completion + + ret = sendXferRange(operation, local, remote, remote_agent, handle, 0, lcnt); + if (ret != NIXL_SUCCESS) { return ret; } - ret = intHandle->status(); + ret = int_handle->status(); if (opt_args && opt_args->hasNotif) { if (ret == NIXL_SUCCESS) { - ret = notifSendPriv(remote_agent, opt_args->notifMsg, req, workerId); - if (_retHelper(ret, intHandle, req)) { + nixlUcxReq req; + auto rmd = (nixlUcxPublicMetadata *)remote[0].metadataP; + ret = notifSendPriv( + remote_agent, opt_args->notifMsg, req, rmd->conn->getEp(int_handle->getWorkerId())); + if (_retHelper(ret, int_handle, req, rmd->conn)) { return ret; } - ret = intHandle->status(); + ret = int_handle->status(); } else if (ret == NIXL_IN_PROG) { - intHandle->notification().emplace(remote_agent, opt_args->notifMsg); + int_handle->notification().emplace(remote_agent, opt_args->notifMsg); } } @@ -1075,22 +1586,33 @@ nixl_status_t nixlUcxEngine::postXfer (const nixl_xfer_op_t &operation, nixl_status_t nixlUcxEngine::checkXfer (nixlBackendReqH* handle) const { nixlUcxBackendH *intHandle = (nixlUcxBackendH *)handle; - size_t workerId = intHandle->getWorkerId(); - - nixl_status_t status = intHandle->status(); auto& notif = intHandle->notification(); - if (status == NIXL_SUCCESS && notif.has_value()) { - nixlUcxReq req; - status = notifSendPriv(notif->agent, notif->payload, req, workerId); - notif.reset(); - if (_retHelper(status, intHandle, req)) { - return status; + nixl_status_t handle_status = intHandle->status(); + + if ((handle_status != NIXL_SUCCESS) || !notif.has_value()) { + if (handle_status != NIXL_IN_PROG) { // error flow + notif.reset(); } - status = intHandle->status(); + return handle_status; } - return status; + ucx_connection_ptr_t conn = getConnection(notif->agent); + if (!conn) { + notif.reset(); + return NIXL_ERR_NOT_FOUND; + } + + nixlUcxReq req; + nixl_status_t status = + notifSendPriv(notif->agent, notif->payload, req, conn->getEp(intHandle->getWorkerId())); + notif.reset(); + status = _retHelper(status, intHandle, req, conn); + if (status != NIXL_SUCCESS) { + return status; + } + + return intHandle->status(); } nixl_status_t nixlUcxEngine::releaseReqH(nixlBackendReqH* handle) const @@ -1104,6 +1626,90 @@ nixl_status_t nixlUcxEngine::releaseReqH(nixlBackendReqH* handle) const return status; } +nixl_status_t +nixlUcxEngine::createGpuXferReq(const nixlBackendReqH &req_hndl, + const nixl_meta_dlist_t &local_descs, + const nixl_meta_dlist_t &remote_descs, + nixlGpuXferReqH &gpu_req_hndl) const { + auto intHandle = static_cast(&req_hndl); + + if (local_descs.descCount() != remote_descs.descCount()) { + NIXL_ERROR << "Mismatch between local and remote descriptor counts"; + return NIXL_ERR_INVALID_PARAM; + } + + if (local_descs.descCount() == 0) { + NIXL_ERROR << "Empty descriptor lists"; + return NIXL_ERR_INVALID_PARAM; + } + + auto remoteMd = static_cast(remote_descs[0].metadataP); + if (!remoteMd || !remoteMd->conn) { + NIXL_ERROR << "No connection found in remote metadata"; + return NIXL_ERR_NOT_FOUND; + } + + size_t workerId = intHandle->getWorkerId(); + nixlUcxEp *ep = remoteMd->conn->getEp(workerId).get(); + + std::vector local_mems; + std::vector remote_rkeys; + local_mems.reserve(local_descs.descCount()); + remote_rkeys.reserve(remote_descs.descCount()); + + for (size_t i = 0; i < static_cast(local_descs.descCount()); i++) { + auto localMd = static_cast(local_descs[i].metadataP); + auto remoteMdDesc = static_cast(remote_descs[i].metadataP); + + local_mems.push_back(localMd->mem); + remote_rkeys.push_back(&remoteMdDesc->getRkey(workerId)); + } + + try { + gpu_req_hndl = nixl::ucx::createGpuXferReq(*ep, local_mems, remote_rkeys); + return NIXL_SUCCESS; + } + catch (const std::exception &e) { + NIXL_ERROR << "Failed to create device memory list for GPU transfer: " << e.what(); + return NIXL_ERR_BACKEND; + } +} + +void +nixlUcxEngine::releaseGpuXferReq(nixlGpuXferReqH gpu_req_hndl) const { + nixl::ucx::releaseGpuXferReq(gpu_req_hndl); +} + +nixl_status_t +nixlUcxEngine::getGpuSignalSize(size_t &signal_size) const { + if (gpuSignalSize_) { + signal_size = *gpuSignalSize_; + return NIXL_SUCCESS; + } + + try { + gpuSignalSize_ = signal_size = uc->getGpuSignalSize(); + return NIXL_SUCCESS; + } + catch (const std::exception &e) { + NIXL_ERROR << e.what(); + return NIXL_ERR_BACKEND; + } +} + +nixl_status_t +nixlUcxEngine::prepGpuSignal(const nixlBackendMD &meta, void *signal) const { + try { + auto *ucx_meta = static_cast(&meta); + getWorker(getWorkerId())->prepGpuSignal(ucx_meta->mem, signal); + return NIXL_SUCCESS; + } + catch (const std::exception &e) { + NIXL_ERROR << e.what(); + return NIXL_ERR_BACKEND; + } +} + int nixlUcxEngine::progress() { // TODO: add listen for connection handling if necessary int ret = 0; @@ -1117,30 +1723,21 @@ int nixlUcxEngine::progress() { *****************************************/ //agent will provide cached msg -nixl_status_t nixlUcxEngine::notifSendPriv(const std::string &remote_agent, - const std::string &msg, - nixlUcxReq &req, - size_t worker_id) const -{ +nixl_status_t +nixlUcxEngine::notifSendPriv(const std::string &remote_agent, + const std::string &msg, + nixlUcxReq &req, + const std::unique_ptr &ep) const { nixlSerDes ser_des; nixl_status_t ret; - auto search = remoteConnMap.find(remote_agent); - - if(search == remoteConnMap.end()) { - //TODO: err: remote connection not found - return NIXL_ERR_NOT_FOUND; - } - ser_des.addStr("name", localAgent); ser_des.addStr("msg", msg); // TODO: replace with mpool for performance - auto buffer = std::make_unique(std::move(ser_des.exportStr())); - ret = search->second->getEp(worker_id)->sendAm(NOTIF_STR, NULL, 0, - (void*)buffer->data(), buffer->size(), - UCP_AM_SEND_FLAG_EAGER, req); - + auto buffer = std::make_unique(ser_des.exportStr()); + ret = ep->sendAm( + NOTIF_STR, NULL, 0, (void *)buffer->data(), buffer->size(), UCP_AM_SEND_FLAG_EAGER, req); if (ret == NIXL_IN_PROG) { nixlUcxIntReq* nReq = (nixlUcxIntReq*)req; nReq->amBuffer = std::move(buffer); @@ -1148,6 +1745,17 @@ nixl_status_t nixlUcxEngine::notifSendPriv(const std::string &remote_agent, return ret; } +ucx_connection_ptr_t +nixlUcxEngine::getConnection(const std::string &remote_agent) const { + auto search = remoteConnMap.find(remote_agent); + return (search != remoteConnMap.end()) ? search->second : nullptr; +} + +void +nixlUcxEngine::appendNotif(std::string remote_name, std::string msg) { + notifMainList.emplace_back(std::move(remote_name), std::move(msg)); +} + ucs_status_t nixlUcxEngine::notifAmCb(void *arg, const void *header, size_t header_length, void *data, @@ -1167,37 +1775,22 @@ nixlUcxEngine::notifAmCb(void *arg, const void *header, std::string remote_name = ser_des.getStr("name"); std::string msg = ser_des.getStr("msg"); - if (engine->isProgressThread()) { - /* Append to the private list to allow batching */ - engine->notifPthrPriv.push_back(std::make_pair(std::move(remote_name), std::move(msg))); - } else { - engine->notifMainList.push_back(std::make_pair(std::move(remote_name), std::move(msg))); - } - + engine->appendNotif(std::move(remote_name), std::move(msg)); return UCS_OK; } -void nixlUcxEngine::notifProgressCombineHelper(notif_list_t &src, notif_list_t &tgt) -{ - const std::lock_guard lock(notifMtx); - moveNotifList(src, tgt); -} - -void nixlUcxEngine::notifProgress() -{ - notifProgressCombineHelper(notifPthrPriv, notifPthr); +void +nixlUcxEngine::getNotifsImpl(notif_list_t ¬if_list) { + moveNotifList(notifMainList, notif_list); } nixl_status_t nixlUcxEngine::getNotifs(notif_list_t ¬if_list) { - if (notif_list.size()!=0) - return NIXL_ERR_INVALID_PARAM; - - if(!pthrOn) while(progress()); - - moveNotifList(notifMainList, notif_list); - notifProgressCombineHelper(notifPthr, notif_list); + if (!notif_list.empty()) return NIXL_ERR_INVALID_PARAM; + while (progress()) + ; + getNotifsImpl(notif_list); return NIXL_SUCCESS; } @@ -1205,14 +1798,17 @@ nixl_status_t nixlUcxEngine::genNotif(const std::string &remote_agent, const std { nixl_status_t ret; nixlUcxReq req; - size_t wid = getWorkerId(); - ret = notifSendPriv(remote_agent, msg, req, wid); + auto conn = getConnection(remote_agent); + if (!conn) { + return NIXL_ERR_NOT_FOUND; + } + ret = notifSendPriv(remote_agent, msg, req, conn->getEp(getWorkerId())); switch(ret) { case NIXL_IN_PROG: /* do not track the request */ - getWorker(wid)->reqRelease(req); + getWorker(getWorkerId())->reqRelease(req); case NIXL_SUCCESS: break; default: diff --git a/src/plugins/ucx/ucx_backend.h b/src/plugins/ucx/ucx_backend.h index c150c13057..51d5ec423f 100644 --- a/src/plugins/ucx/ucx_backend.h +++ b/src/plugins/ucx/ucx_backend.h @@ -27,6 +27,7 @@ #include #include #include +#include #include "nixl.h" #include "backend/backend_engine.h" @@ -36,9 +37,8 @@ #include "common/nixl_time.h" #include "ucx/rkey.h" #include "ucx/ucx_utils.h" -#include "common/list_elem.h" -enum ucx_cb_op_t {CONN_CHECK, NOTIF_STR, DISCONNECT}; +enum ucx_cb_op_t { NOTIF_STR }; class nixlUcxConnection : public nixlBackendConnMD { private: @@ -103,165 +103,287 @@ class nixlUcxCudaCtx; class nixlUcxCudaDevicePrimaryCtx; using nixlUcxCudaDevicePrimaryCtxPtr = std::shared_ptr; -class nixlUcxEngine - : public nixlBackendEngine { - private: - /* UCX data */ - std::unique_ptr uc; - std::vector> uws; - std::string workerAddr; - - /* Progress thread data */ - std::mutex pthrActiveLock; - std::condition_variable pthrActiveCV; - bool pthrActive; - bool pthrOn; - std::thread pthr; - std::chrono::milliseconds pthrDelay; - int pthrControlPipe[2]; - std::vector pollFds; - - /* CUDA data*/ - std::unique_ptr cudaCtx; // Context matching specific device - bool cuda_addr_wa; - - // Context to use when current context is missing - nixlUcxCudaDevicePrimaryCtxPtr m_cudaPrimaryCtx; - - /* Notifications */ - notif_list_t notifMainList; - std::mutex notifMtx; - notif_list_t notifPthrPriv, notifPthr; - - // Map of agent name to saved nixlUcxConnection info - std::unordered_map, strEqual> remoteConnMap; - - - void vramInitCtx(); - void vramFiniCtx(); - int vramUpdateCtx(void *address, uint64_t devId, bool &restart_reqd); - int vramApplyCtx(); - - // Threading infrastructure - // TODO: move the thread management one outside of NIXL common infra - void progressFunc(); - void progressThreadStart(); - void progressThreadStop(); - void progressThreadRestart(); - bool isProgressThread() const noexcept { - return std::this_thread::get_id() == pthr.get_id(); - } +class nixlUcxEngine : public nixlBackendEngine { +public: + static std::unique_ptr + create(const nixlBackendInitParams &init_params); - // Connection helper - static ucs_status_t - connectionCheckAmCb(void *arg, const void *header, - size_t header_length, void *data, - size_t length, - const ucp_am_recv_param_t *param); - - static ucs_status_t - connectionTermAmCb(void *arg, const void *header, - size_t header_length, void *data, - size_t length, - const ucp_am_recv_param_t *param); - - // Memory management helpers - nixl_status_t internalMDHelper (const nixl_blob_t &blob, - const std::string &agent, - nixlBackendMD* &output); - - // Notifications - static ucs_status_t notifAmCb(void *arg, const void *header, - size_t header_length, void *data, - size_t length, - const ucp_am_recv_param_t *param); - nixl_status_t notifSendPriv(const std::string &remote_agent, - const std::string &msg, - nixlUcxReq &req, - size_t worker_id) const; - void notifProgress(); - void notifProgressCombineHelper(notif_list_t &src, notif_list_t &tgt); + ~nixlUcxEngine(); - public: - nixlUcxEngine(const nixlBackendInitParams* init_params); - ~nixlUcxEngine(); - - bool supportsRemote() const override { return true; } - bool supportsLocal() const override { return true; } - bool supportsNotif() const override { return true; } - bool supportsProgTh() const override { return pthrOn; } - - nixl_mem_list_t getSupportedMems() const override; - - /* Object management */ - nixl_status_t getPublicData (const nixlBackendMD* meta, - std::string &str) const override; - nixl_status_t getConnInfo(std::string &str) const override; - nixl_status_t loadRemoteConnInfo (const std::string &remote_agent, - const std::string &remote_conn_info) override; - - nixl_status_t connect(const std::string &remote_agent) override; - nixl_status_t disconnect(const std::string &remote_agent) override; - - nixl_status_t registerMem (const nixlBlobDesc &mem, - const nixl_mem_t &nixl_mem, - nixlBackendMD* &out) override; - nixl_status_t deregisterMem (nixlBackendMD* meta) override; - - nixl_status_t loadLocalMD (nixlBackendMD* input, - nixlBackendMD* &output) override; - - nixl_status_t loadRemoteMD (const nixlBlobDesc &input, - const nixl_mem_t &nixl_mem, - const std::string &remote_agent, - nixlBackendMD* &output) override; - nixl_status_t unloadMD (nixlBackendMD* input) override; - - // Data transfer - nixl_status_t prepXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args=nullptr) const override; - - nixl_status_t estimateXferCost(const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* const &handle, - std::chrono::microseconds &duration, - std::chrono::microseconds &err_margin, - nixl_cost_t &method, - const nixl_opt_args_t* opt_args=nullptr) const override; - - nixl_status_t postXfer (const nixl_xfer_op_t &operation, - const nixl_meta_dlist_t &local, - const nixl_meta_dlist_t &remote, - const std::string &remote_agent, - nixlBackendReqH* &handle, - const nixl_opt_b_args_t* opt_args=nullptr) const override; - - nixl_status_t checkXfer (nixlBackendReqH* handle) const override; - nixl_status_t releaseReqH(nixlBackendReqH* handle) const override; - - int progress() override; - - nixl_status_t getNotifs(notif_list_t ¬if_list); - nixl_status_t genNotif(const std::string &remote_agent, const std::string &msg) const override; - - //public function for UCX worker to mark connections as connected - nixl_status_t checkConn(const std::string &remote_agent); - nixl_status_t endConn(const std::string &remote_agent); - - const std::unique_ptr &getWorker(size_t worker_id) const { - return uws[worker_id]; - } + bool + supportsRemote() const override { + return true; + } - size_t getWorkerId() const { - return std::hash{}(std::this_thread::get_id()) % uws.size(); - } + bool + supportsLocal() const override { + return true; + } + + bool + supportsNotif() const override { + return true; + } + + nixl_mem_list_t + getSupportedMems() const override; + + /* Object management */ + nixl_status_t + getPublicData(const nixlBackendMD *meta, std::string &str) const override; + nixl_status_t + getConnInfo(std::string &str) const override; + nixl_status_t + loadRemoteConnInfo(const std::string &remote_agent, + const std::string &remote_conn_info) override; + + nixl_status_t + connect(const std::string &remote_agent) override; + nixl_status_t + disconnect(const std::string &remote_agent) override; + + nixl_status_t + registerMem(const nixlBlobDesc &mem, const nixl_mem_t &nixl_mem, nixlBackendMD *&out) override; + nixl_status_t + deregisterMem(nixlBackendMD *meta) override; + + nixl_status_t + loadLocalMD(nixlBackendMD *input, nixlBackendMD *&output) override; + + nixl_status_t + loadRemoteMD(const nixlBlobDesc &input, + const nixl_mem_t &nixl_mem, + const std::string &remote_agent, + nixlBackendMD *&output) override; + nixl_status_t + unloadMD(nixlBackendMD *input) override; + + // Data transfer + nixl_status_t + prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const override; + + nixl_status_t + estimateXferCost(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *const &handle, + std::chrono::microseconds &duration, + std::chrono::microseconds &err_margin, + nixl_cost_t &method, + const nixl_opt_args_t *opt_args = nullptr) const override; + + nixl_status_t + postXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const override; + + nixl_status_t + checkXfer(nixlBackendReqH *handle) const override; + nixl_status_t + releaseReqH(nixlBackendReqH *handle) const override; + + nixl_status_t + createGpuXferReq(const nixlBackendReqH &req_hndl, + const nixl_meta_dlist_t &local_descs, + const nixl_meta_dlist_t &remote_descs, + nixlGpuXferReqH &gpu_req_hndl) const override; + + void + releaseGpuXferReq(nixlGpuXferReqH gpu_req_hndl) const override; + + nixl_status_t + getGpuSignalSize(size_t &signal_size) const override; + + nixl_status_t + prepGpuSignal(const nixlBackendMD &meta, void *signal) const override; + + int + progress(); + + nixl_status_t + getNotifs(notif_list_t ¬if_list) override; + nixl_status_t + genNotif(const std::string &remote_agent, const std::string &msg) const override; + + // public function for UCX worker to mark connections as connected + nixl_status_t + checkConn(const std::string &remote_agent); + +protected: + const std::vector> & + getWorkers() const { + return uws; + } + + const std::unique_ptr & + getWorker(size_t worker_id) const { + return uws[worker_id]; + } + + size_t + getWorkerId() const; + + virtual size_t + getSharedWorkersSize() const { + return uws.size(); + } + + void + getNotifsImpl(notif_list_t ¬if_list); + + virtual int + vramApplyCtx(); + + virtual void + appendNotif(std::string remote_name, std::string msg); + + virtual nixl_status_t + sendXferRange(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *handle, + size_t start_idx, + size_t end_idx) const; + + nixlUcxEngine(const nixlBackendInitParams &init_params); + +private: + void + vramInitCtx(); + void + vramFiniCtx(); + int + vramUpdateCtx(void *address, uint64_t dev_id, bool &restart_reqd); + + // Memory management helpers + nixl_status_t + internalMDHelper(const nixl_blob_t &blob, const std::string &agent, nixlBackendMD *&output); + + // Notifications + static ucs_status_t + notifAmCb(void *arg, + const void *header, + size_t header_length, + void *data, + size_t length, + const ucp_am_recv_param_t *param); + + nixl_status_t + notifSendPriv(const std::string &remote_agent, + const std::string &msg, + nixlUcxReq &req, + const std::unique_ptr &ep) const; + + ucx_connection_ptr_t + getConnection(const std::string &remote_agent) const; + + /* UCX data */ + std::unique_ptr uc; + std::vector> uws; + std::string workerAddr; + mutable std::atomic sharedWorkerIndex_; + + /* CUDA data*/ + std::unique_ptr cudaCtx; // Context matching specific device + bool cuda_addr_wa; + mutable std::optional gpuSignalSize_; + + // Context to use when current context is missing + nixlUcxCudaDevicePrimaryCtxPtr m_cudaPrimaryCtx; + + /* Notifications */ + notif_list_t notifMainList; + + // Map of agent name to saved nixlUcxConnection info + std::unordered_map, strEqual> + remoteConnMap; +}; + +class nixlUcxThread; + +/** + * Represents an engine with a single progress thread for all shared workers + */ +class nixlUcxThreadEngine : public nixlUcxEngine { +public: + nixlUcxThreadEngine(const nixlBackendInitParams &init_params); + ~nixlUcxThreadEngine(); + + nixl_status_t + getNotifs(notif_list_t ¬if_list) override; + +protected: + int + vramApplyCtx() override; + + void + appendNotif(std::string remote_name, std::string msg) override; + +private: + std::unique_ptr thread_; + std::mutex notifMtx_; + notif_list_t notifPthr_; +}; + +namespace asio { +class io_context; +} + +class nixlUcxThreadPoolEngine : public nixlUcxEngine { +public: + nixlUcxThreadPoolEngine(const nixlBackendInitParams &init_params); + ~nixlUcxThreadPoolEngine(); + + nixl_status_t + prepXfer(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *&handle, + const nixl_opt_b_args_t *opt_args = nullptr) const override; + + size_t + getSharedWorkersSize() const override { + return numSharedWorkers_; + } + + nixl_status_t + getNotifs(notif_list_t ¬if_list) override; + +protected: + int + vramApplyCtx() override; + + void + appendNotif(std::string remote_name, std::string msg) override; + + nixl_status_t + sendXferRange(const nixl_xfer_op_t &operation, + const nixl_meta_dlist_t &local, + const nixl_meta_dlist_t &remote, + const std::string &remote_agent, + nixlBackendReqH *handle, + size_t start_idx, + size_t end_idx) const override; + +private: + std::unique_ptr io_; + std::unique_ptr sharedThread_; + std::vector> dedicatedThreads_; + size_t numSharedWorkers_; + std::mutex notifMutex_; + notif_list_t notifThread_; + size_t splitBatchSize_; }; #endif diff --git a/src/plugins/ucx/ucx_plugin.cpp b/src/plugins/ucx/ucx_plugin.cpp index 425cf1407b..eb087c7a3e 100644 --- a/src/plugins/ucx/ucx_plugin.cpp +++ b/src/plugins/ucx/ucx_plugin.cpp @@ -18,73 +18,28 @@ #include "backend/backend_plugin.h" #include "ucx_backend.h" -#include "nixl_log.h" - -namespace -{ - const char* ucx_plugin_name = "UCX"; - const char* ucx_plugin_version = "0.1.0"; - - [[nodiscard]] nixlBackendEngine* create_ucx_engine(const nixlBackendInitParams* init_params) { - try { - return new nixlUcxEngine(init_params); - } catch (const std::exception &e) { - NIXL_ERROR << "Failed to create UCX engine: " << e.what(); - return nullptr; - } - } - - void destroy_ucx_engine(nixlBackendEngine *engine) { - delete engine; - } - - [[nodiscard]] const char* get_plugin_name() { - return ucx_plugin_name; - } - - [[nodiscard]] const char* get_plugin_version() { - return ucx_plugin_version; - } - - [[nodiscard]] nixl_b_params_t get_backend_options() { - return get_ucx_backend_common_options(); - } - - [[nodiscard]] nixl_mem_list_t get_backend_mems() { - return { - DRAM_SEG, - VRAM_SEG - }; - } - - // Static plugin structure - nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_ucx_engine, - destroy_ucx_engine, - get_plugin_name, - get_plugin_version, - get_backend_options, - get_backend_mems - }; - -} // namespace +// Plugin type alias for convenience +using ucx_plugin_t = nixlBackendPluginCreator; #ifdef STATIC_PLUGIN_UCX - -nixlBackendPlugin* createStaticUcxPlugin() { - return &plugin; +nixlBackendPlugin * +createStaticUCXPlugin() { + return ucx_plugin_t::create(NIXL_PLUGIN_API_VERSION, + "UCX", + "0.1.0", + get_ucx_backend_common_options(), + {DRAM_SEG, VRAM_SEG}); } - #else - -// Plugin initialization function -extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - return &plugin; +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return ucx_plugin_t::create(NIXL_PLUGIN_API_VERSION, + "UCX", + "0.1.0", + get_ucx_backend_common_options(), + {DRAM_SEG, VRAM_SEG}); } -// Plugin cleanup function -extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { - // Cleanup any resources if needed -} +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} #endif diff --git a/src/plugins/ucx_mo/meson.build b/src/plugins/ucx_mo/meson.build index 2debb97d61..6b4563ea26 100644 --- a/src/plugins/ucx_mo/meson.build +++ b/src/plugins/ucx_mo/meson.build @@ -39,7 +39,9 @@ else install: true, cpp_args : compile_flags + ['-fPIC'], name_prefix: 'libplugin_', # Custom prefix for plugin libraries - install_dir: plugin_install_dir) + install_dir: plugin_install_dir, + # FIXME: normally plugins should not depend directly on each other + install_rpath: '$ORIGIN:$ORIGIN/..') if get_option('buildtype') == 'debug' run_command('sh', '-c', diff --git a/src/plugins/ucx_mo/ucx_mo_backend.cpp b/src/plugins/ucx_mo/ucx_mo_backend.cpp index 5630e9e7a0..c3d3c1368d 100644 --- a/src/plugins/ucx_mo/ucx_mo_backend.cpp +++ b/src/plugins/ucx_mo/ucx_mo_backend.cpp @@ -189,7 +189,7 @@ nixlUcxMoEngine::nixlUcxMoEngine(const nixlBackendInitParams* init_params): setEngCnt(num_ucx_engines); // Initialize required number of engines for (uint32_t i = 0; i < getEngCnt(); i++) { - auto e = std::make_unique(init_params); + auto e = nixlUcxEngine::create(*init_params); if (e->getInitErr()) { this->initErr = true; // TODO: Log error @@ -545,13 +545,9 @@ nixlUcxMoEngine::prepXfer (const nixl_xfer_op_t &operation, /* Allocate internal dlists if needed */ if (!req->dlMatrix[lidx][ridx].in_use) { req->dlMatrix[lidx][ridx].in_use = true; - req->dlMatrix[lidx][ridx].ldescs = new nixl_meta_dlist_t ( - local.getType(), - local.isSorted()); + req->dlMatrix[lidx][ridx].ldescs = new nixl_meta_dlist_t(local.getType()); - req->dlMatrix[lidx][ridx].rdescs = new nixl_meta_dlist_t ( - remote.getType(), - remote.isSorted()); + req->dlMatrix[lidx][ridx].rdescs = new nixl_meta_dlist_t(remote.getType()); } nixlMetaDesc ldesc = local[i]; diff --git a/src/plugins/ucx_mo/ucx_mo_backend.h b/src/plugins/ucx_mo/ucx_mo_backend.h index b93a1c9a25..9fef3a2518 100644 --- a/src/plugins/ucx_mo/ucx_mo_backend.h +++ b/src/plugins/ucx_mo/ucx_mo_backend.h @@ -30,7 +30,6 @@ // Local includes #include -#include #include class nixlUcxMoConnection : public nixlBackendConnMD { @@ -96,7 +95,7 @@ class nixlUcxMoEngine : public nixlBackendEngine { bool pthrOn; // UCX backends data - std::vector> engines; + std::vector> engines; // Map of agent name to saved nixlUcxConnection info using remote_conn_map_t = std::map; using remote_comm_it_t = remote_conn_map_t::iterator; @@ -115,7 +114,6 @@ class nixlUcxMoEngine : public nixlBackendEngine { bool supportsRemote () const { return true; } bool supportsLocal () const { return false; } bool supportsNotif () const { return true; } - bool supportsProgTh () const { return pthrOn; } nixl_mem_list_t getSupportedMems () const; @@ -160,14 +158,16 @@ class nixlUcxMoEngine : public nixlBackendEngine { nixl_status_t checkXfer (nixlBackendReqH* handle) const; nixl_status_t releaseReqH(nixlBackendReqH* handle) const; - int progress(); - nixl_status_t getNotifs(notif_list_t ¬if_list); nixl_status_t genNotif(const std::string &remote_agent, const std::string &msg) const; //public function for UCX worker to mark connections as connected nixl_status_t checkConn(const std::string &remote_agent); nixl_status_t endConn(const std::string &remote_agent); + + // Public function as it is used in tests + int + progress(); }; #endif diff --git a/src/plugins/ucx_mo/ucx_mo_plugin.cpp b/src/plugins/ucx_mo/ucx_mo_plugin.cpp index cc8c8738bd..c5442d5fd7 100644 --- a/src/plugins/ucx_mo/ucx_mo_plugin.cpp +++ b/src/plugins/ucx_mo/ucx_mo_plugin.cpp @@ -15,56 +15,35 @@ * limitations under the License. */ - #include "backend/backend_plugin.h" - #include "ucx_mo_backend.h" - #include "ucx_utils.h" +#include "backend/backend_plugin.h" +#include "ucx_mo_backend.h" +#include "ucx_utils.h" - // Plugin version information - static const char* PLUGIN_NAME = "UCX_MO"; - static const char* PLUGIN_VERSION = "0.1.0"; - // Function to create a new UCX backend engine instance - static nixlBackendEngine* create_engine(const nixlBackendInitParams* init_params) - { - return new nixlUcxMoEngine(init_params); - } - static void destroy_engine(nixlBackendEngine *engine) - { - delete (nixlUcxMoEngine*)engine; - } - // Function to get the plugin name - static const char* get_plugin_name() { - return PLUGIN_NAME; - } - // Function to get the plugin version - static const char* get_plugin_version() { - return PLUGIN_VERSION; - } - // Function to get backend options - static nixl_b_params_t get_backend_options() { - nixl_b_params_t params = get_ucx_backend_common_options(); - params["num_ucx_engines"] = "8"; - return params; - } - // Static plugin structure - static nixlBackendPlugin plugin = { - NIXL_PLUGIN_API_VERSION, - create_engine, - destroy_engine, - get_plugin_name, - get_plugin_version, - get_backend_options - }; - #ifdef STATIC_PLUGIN_UCX_MO - nixlBackendPlugin* createStaticUcxMoPlugin() { - return &plugin; // Return the static plugin instance - } - #else - // Plugin initialization function - extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin* nixl_plugin_init() { - return &plugin; - } - // Plugin cleanup function - extern "C" NIXL_PLUGIN_EXPORT void nixl_plugin_fini() { - // Cleanup any resources if needed - } - #endif +namespace { +nixl_b_params_t +get_ucx_mo_options() { + nixl_b_params_t params = get_ucx_backend_common_options(); + params["num_ucx_engines"] = "8"; + return params; +} +} // namespace + +// Plugin type alias for convenience +using ucx_mo_plugin_t = nixlBackendPluginCreator; + +#ifdef STATIC_PLUGIN_UCX_MO +nixlBackendPlugin * +createStaticUCX_MOPlugin() { + return ucx_mo_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "UCX_MO", "0.1.0", get_ucx_mo_options(), {DRAM_SEG, VRAM_SEG}); +} +#else +extern "C" NIXL_PLUGIN_EXPORT nixlBackendPlugin * +nixl_plugin_init() { + return ucx_mo_plugin_t::create( + NIXL_PLUGIN_API_VERSION, "UCX_MO", "0.1.0", get_ucx_mo_options(), {DRAM_SEG, VRAM_SEG}); +} + +extern "C" NIXL_PLUGIN_EXPORT void +nixl_plugin_fini() {} +#endif diff --git a/src/utils/common/cyclic_buffer.h b/src/utils/common/cyclic_buffer.h new file mode 100644 index 0000000000..c085ce66a3 --- /dev/null +++ b/src/utils/common/cyclic_buffer.h @@ -0,0 +1,77 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef _NIXL_CYCLIC_BUFFER_HPP +#define _NIXL_CYCLIC_BUFFER_HPP + +#include +#include +#include +#include +#include +#include + +template class sharedRingBuffer { +public: + sharedRingBuffer(const std::string &name, bool create, int version, size_t size = 0); + ~sharedRingBuffer(); + + // Non-copyable + sharedRingBuffer(const sharedRingBuffer &) = delete; + sharedRingBuffer & + operator=(const sharedRingBuffer &) = delete; + + bool + push(const T &item); + bool + pop(T &item); + size_t + size() const; + bool + empty() const; + bool + full() const; + uint32_t + version() const; + size_t + capacity() const; + +private: + struct bufferHeader { + std::atomic write_pos{0}; + std::atomic read_pos{0}; + std::atomic version{0}; + int expected_version{0}; + const size_t capacity; + size_t mask; + + bufferHeader(size_t size); + }; + + size_t + getTotalSize() const; + void + createCyclicBuffer(const std::string &name, int version); + void + openCyclicBuffer(const std::string &name, int version); + + bufferHeader *header_; + T *data_; + size_t bufferSize_; +}; + +#include "cyclic_buffer.tpp" +#endif // _NIXL_CYCLIC_BUFFER_HPP diff --git a/src/utils/common/cyclic_buffer.tpp b/src/utils/common/cyclic_buffer.tpp new file mode 100644 index 0000000000..d3071a36ce --- /dev/null +++ b/src/utils/common/cyclic_buffer.tpp @@ -0,0 +1,257 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "cyclic_buffer.h" + +#include +#include +#include +#include + +#include "nixl_log.h" +#include "util.h" + +template +sharedRingBuffer::sharedRingBuffer(const std::string &name, bool create, int version, size_t size) + : header_(nullptr), + data_(nullptr), + bufferSize_(size) { + + if (create) { + createCyclicBuffer(name, version); + } else { + openCyclicBuffer(name, version); + } +} + +template +sharedRingBuffer::~sharedRingBuffer() { + if (header_) { + msync(header_, getTotalSize(), MS_SYNC); + munmap(header_, getTotalSize()); + } +} + +template +bool +sharedRingBuffer::push(const T &item) { + size_t write_pos = header_->write_pos.load(std::memory_order_relaxed); + size_t next_write = (write_pos + 1) & header_->mask; + + if (next_write == header_->read_pos.load(std::memory_order_acquire)) + return false; // Buffer full + + data_[write_pos] = item; + + header_->write_pos.store(next_write, std::memory_order_release); + return true; +} + +template +bool +sharedRingBuffer::pop(T &item) { + size_t read_pos = header_->read_pos.load(std::memory_order_relaxed); + + if (read_pos == header_->write_pos.load(std::memory_order_acquire)) return false; + + // Read data + item = data_[read_pos]; + + // Update read position + size_t next_read = (read_pos + 1) & header_->mask; + header_->read_pos.store(next_read, std::memory_order_release); + return true; +} + +template +size_t +sharedRingBuffer::size() const { + size_t write_pos = header_->write_pos.load(std::memory_order_acquire); + size_t read_pos = header_->read_pos.load(std::memory_order_acquire); + return (write_pos - read_pos) & header_->mask; +} + +template +bool +sharedRingBuffer::empty() const { + return header_->read_pos.load(std::memory_order_acquire) == + header_->write_pos.load(std::memory_order_acquire); +} + +template +bool +sharedRingBuffer::full() const { + size_t write_pos = header_->write_pos.load(std::memory_order_acquire); + size_t next_write = (write_pos + 1) & header_->mask; + return next_write == header_->read_pos.load(std::memory_order_acquire); +} + +template +uint32_t +sharedRingBuffer::version() const { + return header_->version.load(std::memory_order_acquire); +} + +template +size_t +sharedRingBuffer::capacity() const { + return header_->capacity; +} + +template +sharedRingBuffer::bufferHeader::bufferHeader(size_t size) : capacity(size), mask(size - 1) { + if ((size & (size - 1)) != 0) { + throw std::invalid_argument("Telemetry buffer size must be a power of 2"); + } + + static_assert(std::is_trivially_copyable::value, + "T must be trivially copyable for shared memory"); +} + +template +size_t +sharedRingBuffer::getTotalSize() const { + return sizeof(bufferHeader) + sizeof(T) * bufferSize_; +} + +template +void +sharedRingBuffer::createCyclicBuffer(const std::string &name, int version) { + NIXL_INFO << "Creating file-based shared memory on path: " << name + << " with size: " << bufferSize_; + if (bufferSize_ == 0) { + throw std::invalid_argument("Cannot create buffer with size 0"); + } + auto file_closer = [](int *fd) { + close(*fd); + delete fd; + }; + + int fd = open(name.c_str(), O_CREAT | O_RDWR, S_IRUSR | S_IWUSR | S_IRGRP | S_IROTH); + if (fd == -1) { + NIXL_ERROR << "Failed to open a file for shared memory: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("Failed to open a file for shared memory"); + } + + std::unique_ptr file_fd(new int(fd), file_closer); + + if (ftruncate(*file_fd, getTotalSize()) == -1) { + NIXL_ERROR << "Failed to set file size: " << name << " with error: " << strerror(errno); + unlink(name.c_str()); + throw std::runtime_error("Failed to set file size"); + } + + void *ptr = mmap(nullptr, getTotalSize(), PROT_READ | PROT_WRITE, MAP_SHARED, *file_fd, 0); + if (ptr == MAP_FAILED) { + NIXL_ERROR << "Failed to map file memory: " << name + << " with error: " << strerror(errno); + unlink(name.c_str()); + throw std::runtime_error("Failed to map file memory"); + } + + header_ = static_cast(ptr); + data_ = reinterpret_cast(static_cast(ptr) + sizeof(bufferHeader)); + + new (header_) bufferHeader(bufferSize_); + header_->version.store(version, std::memory_order_release); + header_->expected_version = version; +} + +template +void +sharedRingBuffer::openCyclicBuffer(const std::string &name, int version) { + // Use a lambda with custom deleter to auto-close the file descriptor + auto file_closer = [](int *fd) { + close(*fd); + delete fd; + }; + + int fd = open(name.c_str(), O_RDWR); + if (fd == -1) { + NIXL_ERROR << "Failed to open a file for shared memory: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("Failed to open a file for shared memory"); + } + + std::unique_ptr file_fd(new int(fd), file_closer); + + // Check file size before mapping + struct stat st; + if (fstat(*file_fd, &st) == -1) { + NIXL_ERROR << "Failed to get file stats: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("Failed to get file stats"); + } + + if (static_cast(st.st_size) < sizeof(bufferHeader)) { + NIXL_ERROR << "File too small for buffer header: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("File too small for buffer header"); + } + + // First, map just the header to read the size + void *header_ptr = + mmap(nullptr, sizeof(bufferHeader), PROT_READ | PROT_WRITE, MAP_SHARED, *file_fd, 0); + if (header_ptr == MAP_FAILED) { + unlink(name.c_str()); + NIXL_ERROR << "Failed to map header memory: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("Failed to map header memory"); + } + + bufferHeader *temp_header = static_cast(header_ptr); + + // Check version compatibility + int current_version = temp_header->version.load(std::memory_order_acquire); + if (current_version != version) { + munmap(temp_header, sizeof(bufferHeader)); + unlink(name.c_str()); + NIXL_ERROR << "Version mismatch: expected " + std::to_string(version) + ", got " + + std::to_string(current_version); + throw std::runtime_error("Version mismatch: expected " + std::to_string(version) + + ", got " + std::to_string(current_version)); + } + + // Read the buffer size from header + bufferSize_ = temp_header->capacity; + NIXL_INFO << "Reading existing buffer with size: " << bufferSize_; + + // Check if file is large enough for the entire buffer + if (static_cast(st.st_size) < getTotalSize()) { + munmap(temp_header, sizeof(bufferHeader)); + NIXL_ERROR << "File too small for buffer data: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("File too small for buffer data"); + } + + // Unmap the header and remap the entire buffer + munmap(temp_header, sizeof(bufferHeader)); + + // Map the entire buffer + void *ptr = mmap(nullptr, getTotalSize(), PROT_READ | PROT_WRITE, MAP_SHARED, *file_fd, 0); + if (ptr == MAP_FAILED) { + NIXL_ERROR << "Failed to map file memory: " << name + << " with error: " << strerror(errno); + throw std::runtime_error("Failed to map file memory"); + } + + header_ = static_cast(ptr); + data_ = reinterpret_cast(static_cast(ptr) + sizeof(bufferHeader)); +} + +template class sharedRingBuffer; diff --git a/src/utils/common/meson.build b/src/utils/common/meson.build index 51213d071b..4ca430c20b 100644 --- a/src/utils/common/meson.build +++ b/src/utils/common/meson.build @@ -28,6 +28,7 @@ nixl_common_deps = [ absl_status_dep, absl_strings_dep, absl_synchronization_dep, + dependency('asio', required: true), ] # Define a shared library for common utilities diff --git a/src/utils/common/nixl_log.h b/src/utils/common/nixl_log.h index 1b0cb1ef18..be1d739aeb 100644 --- a/src/utils/common/nixl_log.h +++ b/src/utils/common/nixl_log.h @@ -48,6 +48,11 @@ */ #define NIXL_PERROR NIXL_ERROR.WithPerror() +/* + * Like NIXL_ERROR, but prefixed with current function name and a colon + */ +#define NIXL_ERROR_FUNC NIXL_ERROR << __FUNCTION__ << ": " + /* * Logs messages unconditionally (maps to Abseil WARNING level) */ diff --git a/src/utils/common/operators.h b/src/utils/common/operators.h new file mode 100644 index 0000000000..44a9d5fbca --- /dev/null +++ b/src/utils/common/operators.h @@ -0,0 +1,39 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_UTILS_COMMON_OPERATORS_H +#define NIXL_SRC_UTILS_COMMON_OPERATORS_H + +#include + +#include "nixl_types.h" + +inline std::ostream & +operator<<(std::ostream &os, const nixl_mem_t value) { + return os << nixlEnumStrings::memTypeStr(value); +} + +inline std::ostream & +operator<<(std::ostream &os, const nixl_xfer_op_t value) { + return os << nixlEnumStrings::xferOpStr(value); +} + +inline std::ostream & +operator<<(std::ostream &os, const nixl_status_t value) { + return os << nixlEnumStrings::statusStr(value); +} + +#endif diff --git a/src/utils/libfabric/libfabric_common.cpp b/src/utils/libfabric/libfabric_common.cpp new file mode 100644 index 0000000000..69c9b6c862 --- /dev/null +++ b/src/utils/libfabric/libfabric_common.cpp @@ -0,0 +1,132 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric_common.h" +#include "common/nixl_log.h" + +#include +#include +#include +#include + +#include +#include + +namespace LibfabricUtils { + + +std::pair> +getAvailableEfaDevices() { + std::unordered_map> provider_devices_map; + std::vector all_efa_devices; + std::string fabric_name; + struct fi_info *hints, *info; + hints = fi_allocinfo(); + if (!hints) { + NIXL_ERROR << "Failed to allocate fi_info for device discovery"; + return {fabric_name, all_efa_devices}; + } + + // Important to initialize this to allow differentiation between EFA and EFA-Direct + hints->mode = ~0; + + // Set required capabilities - let libfabric select the best provider + hints->caps = FI_READ | FI_WRITE | FI_RECV | FI_SEND | FI_REMOTE_READ | FI_REMOTE_WRITE | + FI_LOCAL_COMM | FI_REMOTE_COMM; + hints->fabric_attr->prov_name = strdup("efa"); + hints->ep_attr->type = FI_EP_RDM; + + int ret = fi_getinfo(FI_VERSION(1, 9), NULL, NULL, 0, hints, &info); + if (ret) { + NIXL_ERROR << "fi_getinfo failed during device discovery: " << fi_strerror(-ret); + fi_freeinfo(hints); + return {fabric_name, all_efa_devices}; + } + + // Process providers and filter for EFA providers with RMA capabilities + for (struct fi_info *cur = info; cur; cur = cur->next) { + if (cur->domain_attr && cur->domain_attr->name && cur->fabric_attr && + cur->fabric_attr->name) { + + std::string device_name = cur->domain_attr->name; + std::string provider_name = cur->fabric_attr->name; + + // Add device to the appropriate provider's vector + provider_devices_map[provider_name].push_back(device_name); + + NIXL_TRACE << "Found EFA device: " << device_name << " with provider: " << provider_name + << " (caps: 0x" << std::hex << cur->caps << std::dec << ")"; + } + } + + fi_freeinfo(info); + fi_freeinfo(hints); + + // Extract device names from the map, prioritizing efa-direct over efa + all_efa_devices.clear(); + if (provider_devices_map.find("efa-direct") != provider_devices_map.end()) { + all_efa_devices = provider_devices_map["efa-direct"]; + fabric_name = "efa-direct"; + NIXL_TRACE << "Using efa-direct provider with " << all_efa_devices.size() << " devices"; + } else if (provider_devices_map.find("efa") != provider_devices_map.end()) { + all_efa_devices = provider_devices_map["efa"]; + fabric_name = "efa"; + NIXL_TRACE << "Using efa provider with " << all_efa_devices.size() << " devices"; + } + + return {fabric_name, all_efa_devices}; +} + +std::string +hexdump(const void *data) { + static constexpr uint HEXDUMP_MAX_LENGTH = 56; + std::stringstream ss; + ss.str().reserve(HEXDUMP_MAX_LENGTH * 3); + const unsigned char *bytes = static_cast(data); + for (size_t i = 0; i < HEXDUMP_MAX_LENGTH; ++i) { + ss << std::hex << std::setw(2) << std::setfill('0') << static_cast(bytes[i]) << " "; + } + return ss.str(); +} + +// Simple counter for pre-allocation only +static uint32_t g_xfer_id_counter = 1; // Start from 1, 0 reserved for special cases + +std::vector +preallocateXferIds(size_t count) { + std::vector xfer_ids; + xfer_ids.reserve(count); + + for (size_t i = 0; i < count; ++i) { + uint32_t xfer_id = g_xfer_id_counter++; + + // Handle wraparound: 20-bit field can hold 0 to 1,048,575 + if (xfer_id > NIXL_XFER_ID_MASK) { + // Reset counter and try again + g_xfer_id_counter = 1; + xfer_id = 1; + g_xfer_id_counter = 2; // Update for next iteration + } + + xfer_ids.push_back(xfer_id); + } + + return xfer_ids; +} + +} // namespace LibfabricUtils diff --git a/src/utils/libfabric/libfabric_common.h b/src/utils/libfabric/libfabric_common.h new file mode 100644 index 0000000000..363fa4c663 --- /dev/null +++ b/src/utils/libfabric/libfabric_common.h @@ -0,0 +1,162 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_COMMON_H +#define NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_COMMON_H + +#include +#include +#include +#include +#include + +#include "nixl.h" + +#include +#include +#include +#include +#include + +// Libfabric configuration constants +#define NIXL_LIBFABRIC_DEFAULT_CONTROL_RAILS 1 +#define NIXL_LIBFABRIC_SEND_RECV_BUFFER_SIZE 8192 +#define NIXL_LIBFABRIC_CQ_SREAD_TIMEOUT_SEC 1 +#define NIXL_LIBFABRIC_DEFAULT_STRIPING_THRESHOLD (128 * 1024) // 128KB +#define NIXL_LIBFABRIC_MAX_XFER_IDS 1024 // Maximum XFER_IDs per notification +#define LF_EP_NAME_MAX_LEN 56 + +// The immediate data associated with an RDMA operation is 32 bits and is divided as follows: +// | 4-bit MSG TYPE flag | 8-bit agent index | 20-bit XFER_ID | + +// Optimized bit field constants (compile-time computed) +#define NIXL_MSG_TYPE_BITS 4 +#define NIXL_AGENT_INDEX_BITS 8 +#define NIXL_XFER_ID_BITS 20 + +// Pre-computed shift amounts for better performance +#define NIXL_MSG_TYPE_SHIFT 0 +#define NIXL_AGENT_INDEX_SHIFT 4 +#define NIXL_XFER_ID_SHIFT 12 + +// Pre-computed masks (compile-time constants) +#define NIXL_MSG_TYPE_MASK 0xFU // 0x0000000F (4 bits) +#define NIXL_AGENT_INDEX_MASK 0xFFU // 0x000000FF (8 bits) +#define NIXL_XFER_ID_MASK 0xFFFFFU // 0x000FFFFF (20 bits) + +// Message type constants +#define NIXL_LIBFABRIC_MSG_CONNECT 0 +#define NIXL_LIBFABRIC_MSG_ACK 1 +#define NIXL_LIBFABRIC_MSG_NOTIFICTION 2 +#define NIXL_LIBFABRIC_MSG_DISCONNECT 3 +#define NIXL_LIBFABRIC_MSG_TRANSFER 4 + +// Single-operation immediate data extraction (no intermediate shifts) +#define NIXL_GET_MSG_TYPE_FROM_IMM(data) ((data) & NIXL_MSG_TYPE_MASK) +#define NIXL_GET_AGENT_INDEX_FROM_IMM(data) \ + (((data) >> NIXL_AGENT_INDEX_SHIFT) & NIXL_AGENT_INDEX_MASK) +#define NIXL_GET_XFER_ID_FROM_IMM(data) (((data) >> NIXL_XFER_ID_SHIFT) & NIXL_XFER_ID_MASK) + +// Single-operation immediate data creation (minimal bit operations) +#define NIXL_MAKE_IMM_DATA(msg_type, agent_idx, xfer_id) \ + (((uint64_t)(msg_type) & NIXL_MSG_TYPE_MASK) | \ + (((uint64_t)(agent_idx) & NIXL_AGENT_INDEX_MASK) << NIXL_AGENT_INDEX_SHIFT) | \ + (((uint64_t)(xfer_id) & NIXL_XFER_ID_MASK) << NIXL_XFER_ID_SHIFT)) + +/** + * @brief Binary notification format to eliminate SerDes string operations + * + * This structure provides a fixed-size, binary format for notifications + * to avoid expensive string serialization/deserialization operations. + * Used for high-performance notification passing between agents. + */ +struct BinaryNotification { + char agent_name[256]; // Fixed-size agent name (null-terminated) + char message[1024]; // Fixed-size message (null-terminated) + uint32_t xfer_id_count; // Number of XFER_IDs + uint32_t + xfer_ids[NIXL_LIBFABRIC_MAX_XFER_IDS]; // Fixed array of XFER_IDs (max 128 per notification) + + /** @brief Clear all fields to zero */ + void + clear() { + memset(this, 0, sizeof(BinaryNotification)); + } + + /** @brief Set agent name with bounds checking */ + void + setAgentName(const std::string &name) { + strncpy(agent_name, name.c_str(), sizeof(agent_name) - 1); + agent_name[sizeof(agent_name) - 1] = '\0'; + } + + /** @brief Set message with bounds checking */ + void + setMessage(const std::string &msg) { + strncpy(message, msg.c_str(), sizeof(message) - 1); + message[sizeof(message) - 1] = '\0'; + } + + /** @brief Add XFER_ID if space available */ + void + addXferId(uint32_t xfer_id) { + if (xfer_id_count < NIXL_LIBFABRIC_MAX_XFER_IDS) { + xfer_ids[xfer_id_count++] = xfer_id; + } + } + + /** @brief Get agent name as string */ + std::string + getAgentName() const { + return std::string(agent_name); + } + + /** @brief Get message as string */ + std::string + getMessage() const { + return std::string(message); + } + + /** @brief Get all XFER_IDs as unordered set */ + std::unordered_set + getXferIds() const { + std::unordered_set result; + for (uint32_t i = 0; i < xfer_id_count; ++i) { + result.insert(xfer_ids[i]); + } + return result; + } +}; + +// Global XFER_ID management +namespace LibfabricUtils { +// Pre-allocate XFER_IDs during initialization (NOT fast path) +std::vector +preallocateXferIds(size_t count); +} // namespace LibfabricUtils + +// Utility functions +namespace LibfabricUtils { +// Device discovery +std::pair> +getAvailableEfaDevices(); +// String utilities +std::string +hexdump(const void *data); +} // namespace LibfabricUtils + +#endif // NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_COMMON_H diff --git a/src/utils/libfabric/libfabric_rail.cpp b/src/utils/libfabric/libfabric_rail.cpp new file mode 100644 index 0000000000..0e382f3c93 --- /dev/null +++ b/src/utils/libfabric/libfabric_rail.cpp @@ -0,0 +1,1078 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric_rail.h" +#include "common/nixl_log.h" +#include "serdes/serdes.h" + +#include +#include + +// RequestPool Base Class Implementation + +RequestPool::RequestPool(size_t pool_size, size_t rail_id) : rail_id_(rail_id) { + requests_.resize(pool_size); + + for (size_t i = 0; i < pool_size; ++i) { + requests_[i].rail_id = rail_id; + requests_[i].in_use = false; + free_indices_.push(i); + } +} + +void +RequestPool::release(nixlLibfabricReq *req) { + if (!req) return; + + std::lock_guard lock(pool_mutex_); + + req->in_use = false; + req->chunk_offset = 0; + req->chunk_size = 0; + req->completion_callback = nullptr; + memset(&req->ctx, 0, sizeof(fi_context2)); + size_t idx = req - &requests_[0]; + free_indices_.push(idx); +} + +nixlLibfabricReq * +RequestPool::findByContext(void *context) const { + std::lock_guard lock(pool_mutex_); + return reinterpret_cast(context); +} + +size_t +RequestPool::getActiveRequestCount() const { + std::lock_guard lock(pool_mutex_); + return requests_.size() - free_indices_.size(); +} + +// ControlRequestPool Implementation + +ControlRequestPool::ControlRequestPool(size_t pool_size, size_t rail_id) + : RequestPool(pool_size, rail_id), + buffer_chunk_(nullptr), + buffer_chunk_size_(0), + buffer_mr_(nullptr) {} + +ControlRequestPool::~ControlRequestPool() { + // Cleanup should have been called explicitly before domain destruction + // This is just a safety check + cleanup(); +} + +void +ControlRequestPool::cleanup() { + if (buffer_mr_) { + fi_close(&buffer_mr_->fid); + buffer_mr_ = nullptr; + } + if (buffer_chunk_) { + free(buffer_chunk_); + buffer_chunk_ = nullptr; + } +} + +nixl_status_t +ControlRequestPool::initializeWithBuffersAndXferIds(struct fid_domain *domain, + const std::vector &xfer_ids) { + if (xfer_ids.size() != requests_.size()) { + return NIXL_ERR_INVALID_PARAM; + } + + // Allocate buffer chunk + buffer_chunk_size_ = BUFFER_SIZE * requests_.size(); + + buffer_chunk_ = malloc(buffer_chunk_size_); + if (!buffer_chunk_) { + NIXL_ERROR << "Standard allocation failed for control request pool on rail " << rail_id_; + return NIXL_ERR_BACKEND; + } + + NIXL_DEBUG << "Allocated " << buffer_chunk_size_ << " bytes for control request pool on rail " + << rail_id_; + + // Register buffer chunk with libfabric + int ret = fi_mr_reg( + domain, buffer_chunk_, buffer_chunk_size_, FI_SEND | FI_RECV, 0, 0, 0, &buffer_mr_, NULL); + if (ret) { + free(buffer_chunk_); + buffer_chunk_ = nullptr; + return NIXL_ERR_BACKEND; + } + // Pre-assign buffers and XFER_IDs to requests + for (size_t i = 0; i < requests_.size(); ++i) { + requests_[i].xfer_id = xfer_ids[i]; + requests_[i].buffer = static_cast(buffer_chunk_) + (i * BUFFER_SIZE); + requests_[i].mr = buffer_mr_; + requests_[i].buffer_size = BUFFER_SIZE; + requests_[i].operation_type = nixlLibfabricReq::SEND; // Default for control + } + return NIXL_SUCCESS; +} + +nixlLibfabricReq * +ControlRequestPool::allocate(size_t needed_size) { + std::lock_guard lock(pool_mutex_); + + if (free_indices_.empty()) { + return nullptr; // No free requests + } + if (needed_size > BUFFER_SIZE) { + return nullptr; // Size too large + } + + size_t idx = free_indices_.top(); + free_indices_.pop(); + + nixlLibfabricReq *req = &requests_[idx]; + req->in_use = true; + + return req; +} + +// DataRequestPool Implementation + +DataRequestPool::DataRequestPool(size_t pool_size, size_t rail_id) + : RequestPool(pool_size, rail_id) {} + +nixl_status_t +DataRequestPool::initializeWithXferIds(const std::vector &xfer_ids) { + if (xfer_ids.size() != requests_.size()) { + return NIXL_ERR_INVALID_PARAM; + } + // Pre-assign XFER_IDs to requests + for (size_t i = 0; i < requests_.size(); ++i) { + requests_[i].xfer_id = xfer_ids[i]; + requests_[i].buffer = nullptr; // No buffers for data requests + requests_[i].mr = nullptr; + requests_[i].buffer_size = 0; + requests_[i].operation_type = nixlLibfabricReq::WRITE; // Default for data + } + return NIXL_SUCCESS; +} + +nixlLibfabricReq * +DataRequestPool::allocate(nixlLibfabricReq::OpType op_type) { + std::lock_guard lock(pool_mutex_); + + if (free_indices_.empty()) { + return nullptr; + } + + size_t idx = free_indices_.top(); + free_indices_.pop(); + + nixlLibfabricReq *req = &requests_[idx]; + req->in_use = true; + req->operation_type = op_type; + + return req; +} + +// Rail Class Implementation + +nixlLibfabricRail::nixlLibfabricRail(const std::string &device, + const std::string &provider, + uint16_t id) + : rail_id(id), + device_name(device), + blocking_cq_sread_supported(true), + control_request_pool_(CONTROL_REQUESTS_PER_RAIL, id), + data_request_pool_(DATA_REQUESTS_PER_RAIL, id) { + // Initialize all pointers to nullptr + info = nullptr; + fabric = nullptr; + domain = nullptr; + endpoint = nullptr; + cq = nullptr; + av = nullptr; + memset(ep_name, 0, sizeof(ep_name)); + + // Initialize all Libfabric resources for this rail + NIXL_TRACE << "Initializing rail " << rail_id << " with device: " << device_name + << ", provider: " << provider; + + // Initialize hints for this rail + struct fi_info *hints = fi_allocinfo(); + if (!hints) { + NIXL_ERROR << "fi_allocinfo failed for rail " << rail_id; + throw std::runtime_error("Failed to allocate fi_info for rail " + std::to_string(rail_id)); + } + hints->caps = 0; + hints->caps = FI_MSG | FI_RMA; + hints->caps |= FI_LOCAL_COMM | FI_REMOTE_COMM; + if (provider.c_str() == std::string("efa-direct")) { + hints->mode = FI_CONTEXT | FI_CONTEXT2; + } else { + hints->mode = FI_CONTEXT; + } + hints->ep_attr->type = FI_EP_RDM; + hints->domain_attr->mr_mode = + FI_MR_LOCAL | FI_MR_HMEM | FI_MR_VIRT_ADDR | FI_MR_ALLOCATED | FI_MR_PROV_KEY; + hints->domain_attr->mr_key_size = 2; + hints->domain_attr->name = strdup(device_name.c_str()); + hints->domain_attr->threading = FI_THREAD_SAFE; + try { + // Get fabric info for this specific device + int ret = fi_getinfo(FI_VERSION(1, 9), NULL, NULL, 0, hints, &info); + if (ret) { + NIXL_ERROR << "fi_getinfo failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_getinfo failed for rail " + std::to_string(rail_id)); + } + + // Create fabric for this rail + ret = fi_fabric(info->fabric_attr, &fabric, NULL); + if (ret) { + NIXL_ERROR << "fi_fabric failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_fabric failed for rail " + std::to_string(rail_id)); + } + NIXL_TRACE << "fabric_attr->name " << info->fabric_attr->name; + // Create domain for this rail + ret = fi_domain(fabric, info, &domain, NULL); + if (ret) { + NIXL_ERROR << "fi_domain failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_domain failed for rail " + std::to_string(rail_id)); + } + + // Create CQ for this rail + struct fi_cq_attr cq_attr = {}; + cq_attr.format = FI_CQ_FORMAT_DATA; + cq_attr.wait_obj = FI_WAIT_UNSPEC; + cq_attr.size = 12288; + ret = fi_cq_open(domain, &cq_attr, &cq, NULL); + if (ret) { + NIXL_WARN << "fi_cq_open failed for rail " << rail_id << ": " << fi_strerror(-ret) + << " - trying FI_WAIT_NONE"; + if (ret == -FI_ENOSYS) { + NIXL_TRACE << "FI_WAIT_UNSPEC not supported, falling back to FI_WAIT_NONE for rail " + << rail_id; + blocking_cq_sread_supported = false; + // If fi_cq_open fails due to FI_WAIT_UNSPEC not supported, we fall back to + // FI_WAIT_NONE and use fi_cq_read in control rails + cq_attr.wait_obj = FI_WAIT_NONE; + ret = fi_cq_open(domain, &cq_attr, &cq, NULL); + if (ret) { + NIXL_ERROR << "fi_cq_open with FI_WAIT_NONE failed for rail " << rail_id << ": " + << fi_strerror(-ret); + throw std::runtime_error("fi_cq_open with FI_WAIT_NONE failed for rail " + + std::to_string(rail_id)); + } + NIXL_TRACE << "fi_cq_open with FI_WAIT_NONE succeeded for rail " << rail_id; + } else { + throw std::runtime_error("fi_cq_open failed for rail " + std::to_string(rail_id)); + } + } + // Verify CQ was properly created + if (!cq) { + NIXL_ERROR << "CQ is null after fi_cq_open for rail " << rail_id; + throw std::runtime_error("CQ creation returned success but pointer is null for rail " + + std::to_string(rail_id)); + } + // Create AV for this rail + struct fi_av_attr av_attr = {}; + av_attr.type = FI_AV_TABLE; + av_attr.count = 1024; + ret = fi_av_open(domain, &av_attr, &av, NULL); + if (ret) { + NIXL_ERROR << "fi_av_open failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_av_open failed for rail " + std::to_string(rail_id)); + } + + // Create endpoint for this rail + ret = fi_endpoint(domain, info, &endpoint, NULL); + if (ret) { + NIXL_ERROR << "fi_endpoint failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_endpoint failed for rail " + std::to_string(rail_id)); + } + + // Bind endpoint with CQ and AV for this rail + ret = fi_ep_bind(endpoint, &cq->fid, FI_TRANSMIT | FI_RECV); + if (ret) { + NIXL_ERROR << "fi_ep_bind cq failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_ep_bind cq failed for rail " + std::to_string(rail_id)); + } + + ret = fi_ep_bind(endpoint, &av->fid, 0); + if (ret) { + NIXL_ERROR << "fi_ep_bind av failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_ep_bind av failed for rail " + std::to_string(rail_id)); + } + + // Disable shared memory transfers for EFA provider to fix same-agent transfers + bool optval = false; + ret = fi_setopt(&endpoint->fid, + FI_OPT_ENDPOINT, + FI_OPT_SHARED_MEMORY_PERMITTED, + &optval, + sizeof(optval)); + if (ret && ret != -FI_ENOSYS) { + NIXL_WARN << "fi_setopt FI_OPT_SHARED_MEMORY_PERMITTED failed for rail " << rail_id + << ": " << fi_strerror(-ret) << " - continuing anyway"; + } else if (ret == 0) { + NIXL_DEBUG << "Successfully disabled shared memory transfers for rail " << rail_id; + } + + // Enable endpoint for this rail + ret = fi_enable(endpoint); + if (ret) { + NIXL_ERROR << "fi_enable failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_enable failed for rail " + std::to_string(rail_id)); + } + + // Get endpoint name for this rail + size_t ep_name_len = sizeof(ep_name); + ret = fi_getname(&endpoint->fid, ep_name, &ep_name_len); + if (ret) { + NIXL_ERROR << "fi_getname failed for rail " << rail_id << ": " << fi_strerror(-ret); + throw std::runtime_error("fi_getname failed for rail " + std::to_string(rail_id)); + } + + // Pre-allocate XFER_IDs for both pools + std::vector control_xfer_ids = + LibfabricUtils::preallocateXferIds(CONTROL_REQUESTS_PER_RAIL); + std::vector data_xfer_ids = + LibfabricUtils::preallocateXferIds(DATA_REQUESTS_PER_RAIL); + + // Initialize control request pool with buffers and XFER_IDs + nixl_status_t status = + control_request_pool_.initializeWithBuffersAndXferIds(domain, control_xfer_ids); + if (status != NIXL_SUCCESS) { + throw std::runtime_error("Failed to initialize control request pool for rail " + + std::to_string(rail_id)); + } + // Initialize data request pool with XFER_IDs only + status = data_request_pool_.initializeWithXferIds(data_xfer_ids); + if (status != NIXL_SUCCESS) { + throw std::runtime_error("Failed to initialize data request pool for rail " + + std::to_string(rail_id)); + } + + NIXL_TRACE << "Initialized request pools: " << CONTROL_REQUESTS_PER_RAIL + << " control requests, " << DATA_REQUESTS_PER_RAIL << " data requests for rail " + << rail_id; + + // Post initial receive using new resource management system + nixlLibfabricReq *recv_req = allocateControlRequest(NIXL_LIBFABRIC_SEND_RECV_BUFFER_SIZE); + if (!recv_req) { + NIXL_ERROR << "Failed to allocate request for initial receive on rail " << rail_id; + throw std::runtime_error("Failed to allocate request for initial receive on rail " + + std::to_string(rail_id)); + } + status = postRecv(recv_req); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to post initial receive on rail " << rail_id; + releaseRequest(recv_req); + throw std::runtime_error("Failed to post initial receive on rail " + + std::to_string(rail_id)); + } + NIXL_TRACE << "Successfully initialized rail " << rail_id; + } + catch (...) { + fi_freeinfo(hints); + throw; + } + fi_freeinfo(hints); +} + +nixlLibfabricRail::~nixlLibfabricRail() { + cleanup(); +} + +bool +nixlLibfabricRail::isProperlyInitialized() const { + return (cq != nullptr && endpoint != nullptr && domain != nullptr); +} + +void +nixlLibfabricRail::cleanup() { + NIXL_TRACE << "Starting cleanup for rail " << rail_id; + // STEP 1: Close endpoint first (it depends on CQ and AV) + if (endpoint) { + NIXL_TRACE << "Closing endpoint for rail " << rail_id; + int ret = fi_close(&endpoint->fid); + if (ret) { + NIXL_WARN << "fi_close endpoint failed for rail " << rail_id << ": " + << fi_strerror(-ret); + } + endpoint = nullptr; + } + // STEP 2: Close CQ and AV (they depend on domain) + if (cq) { + NIXL_TRACE << "Closing completion queue for rail " << rail_id; + int ret = fi_close(&cq->fid); + if (ret) { + NIXL_WARN << "fi_close cq failed for rail " << rail_id << ": " << fi_strerror(-ret); + } + cq = nullptr; + } + if (av) { + NIXL_TRACE << "Closing address vector for rail " << rail_id; + int ret = fi_close(&av->fid); + if (ret) { + NIXL_WARN << "fi_close av failed for rail " << rail_id << ": " << fi_strerror(-ret); + } + av = nullptr; + } + + // STEP 3: Clean up request pools while domain is still valid + // This ensures all memory registrations (MRs) are properly deregistered before domain closure + NIXL_TRACE << "Cleaning up request pools for rail " << rail_id; + control_request_pool_.cleanup(); + // STEP 4: Close domain AFTER all MRs, endpoint, CQ, AV are closed + if (domain) { + NIXL_TRACE << "Closing domain for rail " << rail_id; + int ret = fi_close(&domain->fid); + if (ret) { + NIXL_WARN << "fi_close domain failed for rail " << rail_id << ": " << fi_strerror(-ret); + } + domain = nullptr; + } + // STEP 5: Close fabric + if (fabric) { + NIXL_TRACE << "Closing fabric for rail " << rail_id; + int ret = fi_close(&fabric->fid); + if (ret) { + NIXL_WARN << "fi_close fabric failed for rail " << rail_id << ": " << fi_strerror(-ret); + } + fabric = nullptr; + } + // STEP 6: Free info structure + if (info) { + NIXL_INFO << "Freeing info structure for rail " << rail_id; + fi_freeinfo(info); + info = nullptr; + } + NIXL_TRACE << "Cleanup completed for rail " << rail_id; +} + +void +nixlLibfabricRail::setNotificationCallback(std::function callback) { + notificationCallback = callback; +} + +void +nixlLibfabricRail::setConnectionAckCallback( + std::function callback) { + connectionAckCallback = callback; +} + +void +nixlLibfabricRail::setConnectionReqCallback( + std::function callback) { + connectionReqCallback = callback; +} + +void +nixlLibfabricRail::setXferIdCallback(std::function callback) { + xferIdCallback = callback; +} + +// Per-Rail Completion Processing + +// Per-rail completion processing - handles one rail's CQ with configurable blocking behavior +nixl_status_t +nixlLibfabricRail::progressCompletionQueue(bool use_blocking) { + // Completion processing + struct fi_cq_data_entry completion; + memset(&completion, 0, sizeof(completion)); + + int ret; + + // Only protect libfabric CQ hardware operations + { + std::lock_guard cq_lock(cq_progress_mutex_); + + if (use_blocking && blocking_cq_sread_supported) { + // Blocking read using fi_cq_sread (used by CM thread) + ret = fi_cq_sread(cq, &completion, 1, nullptr, NIXL_LIBFABRIC_CQ_SREAD_TIMEOUT_SEC); + } else { + // Non-blocking read (used by progress thread or fallback) + ret = fi_cq_read(cq, &completion, 1); + } + + if (ret < 0 && ret != -FI_EAGAIN) { + NIXL_ERROR << "fi_cq_read returned error " << ret << " on rail " << rail_id << ": " + << fi_strerror(-ret); + + // Handle error - but be careful about fi_cq_readerr + struct fi_cq_err_entry err_entry; + memset(&err_entry, 0, sizeof(err_entry)); + + int err_ret = fi_cq_readerr(cq, &err_entry, 0); + if (err_ret > 0) { + NIXL_ERROR << "CQ read failed on rail " << rail_id + << " with error: " << fi_strerror(err_entry.err) + << " prov_errno: " << err_entry.prov_errno << " len: " << err_entry.len; + } else { + NIXL_ERROR << "fi_cq_readerr failed with " << err_ret; + } + return NIXL_ERR_BACKEND; + } + } + // CQ lock released here - completion is now local data + + if (ret == -FI_EAGAIN) { + return NIXL_IN_PROG; // No completions available + } + + if (ret == 1) { + NIXL_TRACE << "Completion received on rail " << rail_id << " flags: " << std::hex + << completion.flags << " data: " << completion.data + << " context: " << completion.op_context << std::dec; + + // Process completion using local data. Callbacks have their own thread safety + nixl_status_t status = processCompletionQueueEntry(&completion); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to process completion on rail " << rail_id; + return status; + } + + NIXL_DEBUG << "Completion processed on rail " << rail_id; + return NIXL_SUCCESS; + } + + return NIXL_ERR_BACKEND; // Unexpected case +} + +// Route completion to appropriate handler (rail-specific) +nixl_status_t +nixlLibfabricRail::processCompletionQueueEntry(struct fi_cq_data_entry *comp) { + uint64_t flags = comp->flags; + + NIXL_TRACE << "Routing completion from rail " << rail_id << " with flags: " << std::hex << flags + << " FI_SEND: " << (flags & FI_SEND) << " FI_RECV: " << (flags & FI_RECV) + << " FI_WRITE: " << (flags & FI_WRITE) + << " FI_REMOTE_WRITE: " << (flags & FI_REMOTE_WRITE) << std::dec; + + if (flags & FI_SEND) { + // Local send completions (fi_senddata) - use context + return processLocalSendCompletion(comp); + + } else if (flags & FI_RECV) { + // Receive completions - use immediate data + return processRecvCompletion(comp); + + } else if (flags & FI_WRITE) { + // Local write completions (fi_writedata) - use context + return processLocalTransferCompletion(comp, "write"); + + } else if (flags & FI_READ) { + // Local read completions (fi_readdata) - use context + return processLocalTransferCompletion(comp, "read"); + + } else if (flags & FI_REMOTE_WRITE) { + // Remote write completions (from fi_writedata) - use immediate data + return processRemoteWriteCompletion(comp); + + } else { + // Add more detailed warning for unknown completion flags + NIXL_WARN << "Unknown completion flags detected on rail " << rail_id << " - flags: 0x" + << std::hex << flags << std::dec << " (FI_SEND=" << !!(flags & FI_SEND) + << " FI_RECV=" << !!(flags & FI_RECV) << " FI_WRITE=" << !!(flags & FI_WRITE) + << " FI_READ=" << !!(flags & FI_READ) + << " FI_REMOTE_WRITE=" << !!(flags & FI_REMOTE_WRITE) + << " FI_REMOTE_READ=" << !!(flags & FI_REMOTE_READ) << ")" + << " data: 0x" << std::hex << comp->data << std::dec + << " context: " << comp->op_context << " len: " << comp->len; + + // Try to find the request associated with this context for debugging + nixlLibfabricReq *req = findRequestFromContext(comp->op_context); + if (req) { + NIXL_WARN << "Found request for zero-flags completion: XFER_ID=" << req->xfer_id + << " context=" << &req->ctx << " req_ptr=" << req << " rail=" << rail_id + << " in_use=" << req->in_use << " op_type=" + << (req->operation_type == nixlLibfabricReq::WRITE ? "WRITE" : + req->operation_type == nixlLibfabricReq::READ ? "READ" : + req->operation_type == nixlLibfabricReq::SEND ? "SEND" : + "RECV"); + } else { + NIXL_WARN << "No request found for zero-flags completion context " << comp->op_context; + } + + // Check if this might be a spurious completion with flags=0 + if (flags == 0) { + NIXL_WARN << "Completion with zero flags detected - this may be a spurious completion " + "or cleanup event"; + // Don't treat zero flags as a fatal error, just skip processing + return NIXL_SUCCESS; + } + + NIXL_ERROR << "Unknown completion flags: " << std::hex << flags << " data: " << comp->data + << " context: " << comp->op_context; + return NIXL_ERR_BACKEND; + } +} + +// Handle local send completions (establishConnection, genNotif) +nixl_status_t +nixlLibfabricRail::processLocalSendCompletion(struct fi_cq_data_entry *comp) { + // Release request back to pool + nixlLibfabricReq *req = findRequestFromContext(comp->op_context); + if (req) { + if (req->completion_callback) { + NIXL_TRACE << "Calling completion callback for send request " << req->xfer_id; + req->completion_callback(); + NIXL_TRACE << "Completion callback completed for send"; + } + // Always release request back to pool + releaseRequest(req); + } else { + NIXL_ERROR << "No request found for context " << comp->op_context << " on rail " << rail_id; + } + + return NIXL_SUCCESS; +} + +// Handle local transfer completions (both read and write operations from postXfer) +nixl_status_t +nixlLibfabricRail::processLocalTransferCompletion(struct fi_cq_data_entry *comp, + const char *operation_type) { + // Find the request from context to access the completion callback + nixlLibfabricReq *req = findRequestFromContext(comp->op_context); + if (req) { + // Call completion callback if it exists + if (req->completion_callback) { + NIXL_TRACE << "Calling completion callback for " << operation_type << " request " + << req->xfer_id; + req->completion_callback(); + NIXL_TRACE << "Completion callback completed for " << operation_type; + } + + // Always release request back to pool + releaseRequest(req); + } else { + NIXL_ERROR << "No request found for " << operation_type << " completion context " + << comp->op_context << " on rail " << rail_id; + } + + return NIXL_SUCCESS; +} + +// Handle remote receive completions (conn_req, conn_ack, notification messages) +nixl_status_t +nixlLibfabricRail::processRecvCompletion(struct fi_cq_data_entry *comp) { + // Get the request from context to access the received buffer + nixlLibfabricReq *req = findRequestFromContext(comp->op_context); + if (!req) { + NIXL_ERROR << "No request found for receive completion context on rail " << rail_id; + return NIXL_ERR_BACKEND; + } + // Decode the immediate data format + uint64_t msg_type = NIXL_GET_MSG_TYPE_FROM_IMM(comp->data); + uint16_t agent_idx = NIXL_GET_AGENT_INDEX_FROM_IMM(comp->data); + uint32_t xfer_id = NIXL_GET_XFER_ID_FROM_IMM(comp->data); + NIXL_TRACE << "Received control message type " << msg_type << " agent_idx=" << agent_idx + << " XFER_ID=" << xfer_id << " imm_data=0x" << std::hex << comp->data << std::dec; + + if (msg_type == NIXL_LIBFABRIC_MSG_CONNECT) { + NIXL_TRACE << "Processing connection request on rail " << rail_id + << " Xfer_id :" << xfer_id; + // Use callback to handle connection request processing + if (connectionReqCallback) { + std::string serialized_data(static_cast(req->buffer), req->buffer_size); + nixl_status_t callback_status = connectionReqCallback( + agent_idx, serialized_data, const_cast(this)); + if (callback_status != NIXL_SUCCESS) { + NIXL_ERROR << "Connection request callback failed"; + return callback_status; + } + NIXL_TRACE << "Connection request processed via callback for rail " << rail_id; + } else { + NIXL_ERROR << "No connection request callback set for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + } else if (msg_type == NIXL_LIBFABRIC_MSG_ACK) { + NIXL_TRACE << "Processing connect request acknowledgement on rail " << rail_id; + // Notify engine that connection is established via callback + // TODO: validate the current state before calling callback + if (connectionAckCallback) { + connectionAckCallback(agent_idx, nullptr, ConnectionState::CONNECTED); + NIXL_TRACE << "Connection state updated to CONNECTED via callback for rail " << rail_id; + } else { + NIXL_ERROR << "No connection state callback set for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + } else if (msg_type == NIXL_LIBFABRIC_MSG_NOTIFICTION) { + NIXL_TRACE << "Processing notification request on rail " << rail_id + << " Xfer_id :" << xfer_id; + + // Create string from received buffer using the actual received length from completion entry + std::string message(static_cast(req->buffer), comp->len); + + NIXL_TRACE << "Adding message: " << message << " to the notification list on rail " + << rail_id; + + // Call engine's callback to store notification in central storage (like reference) + if (notificationCallback) { + notificationCallback(message); + NIXL_TRACE << "Notification stored via callback"; + } else { + NIXL_ERROR << "No notification callback set!"; + return NIXL_ERR_BACKEND; + } + } else if (msg_type == NIXL_LIBFABRIC_MSG_DISCONNECT) { + NIXL_TRACE << "Processing disconnect request from agent " << agent_idx << " on rail " + << rail_id + << "Currently not tracking the fi_addrs, so no callback for disconnect to clean " + "up libfabric AV list"; + } else { + NIXL_ERROR << "Unknown message type: " << std::hex << msg_type << std::dec; + return NIXL_ERR_BACKEND; + } + + // Clear the receive buffer after processing + memset(req->buffer, 0, req->buffer_size); + + // Release the current request + releaseRequest(req); + + // Post a new receive using new resource management system + nixlLibfabricReq *new_req = allocateControlRequest(NIXL_LIBFABRIC_SEND_RECV_BUFFER_SIZE); + if (!new_req) { + NIXL_ERROR << "Failed to allocate request for subsequent receive on rail " << rail_id; + return NIXL_ERR_BACKEND; + } + nixl_status_t status = postRecv(new_req); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to post subsequent receive on rail " << rail_id; + releaseRequest(new_req); + return status; + } + return NIXL_SUCCESS; +} + +// Handle remote write completions (data arrival notification) +nixl_status_t +nixlLibfabricRail::processRemoteWriteCompletion(struct fi_cq_data_entry *comp) const { + // Decode the immediate data format + uint64_t msg_type = NIXL_GET_MSG_TYPE_FROM_IMM(comp->data); + uint16_t agent_idx = NIXL_GET_AGENT_INDEX_FROM_IMM(comp->data); + uint32_t xfer_id = NIXL_GET_XFER_ID_FROM_IMM(comp->data); + + // For remote write completions, we don't need to post a new receive + // The write operation doesn't consume a receive buffer + if (msg_type == NIXL_LIBFABRIC_MSG_TRANSFER) { + NIXL_TRACE << "Remote write completion on rail " << rail_id << " - received " << comp->len + << " bytes" << " agent_idx=" << agent_idx << " XFER_ID=" << xfer_id + << " imm_data=0x" << std::hex << comp->data << std::dec; + + // Call XFER_ID tracking callback to add received XFER_ID to global set + if (xferIdCallback) { + xferIdCallback(xfer_id); + NIXL_TRACE << "Called XFER_ID callback for XFER_ID " << xfer_id; + } else { + NIXL_ERROR << "No XFER_ID callback set for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + } + return NIXL_SUCCESS; +} + +// Per-Rail Libfabric Operation Wrappers + +nixl_status_t +nixlLibfabricRail::postRecv(nixlLibfabricReq *req) const { + if (!req || !req->buffer) { + NIXL_ERROR << "Invalid request or buffer for receive on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + + struct fi_msg msg = {}; + struct iovec msg_iov; + void *desc = fi_mr_desc(req->mr); // Use request's MR + + // Setup message structure using request's buffer + msg_iov.iov_base = req->buffer; + msg_iov.iov_len = req->buffer_size; + + msg.msg_iov = &msg_iov; + msg.desc = &desc; + msg.iov_count = 1; + msg.addr = FI_ADDR_UNSPEC; + msg.context = &req->ctx; // Use request's context directly + msg.data = 0; + + NIXL_TRACE << "Posting receive on endpoint: " << endpoint << " buffer: " << req->buffer + << " size: " << req->buffer_size << " context: " << &req->ctx; + + int ret = fi_recvmsg(endpoint, &msg, 0); + if (ret) { + NIXL_ERROR << "fi_recvmsg failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + + NIXL_TRACE << "Receive posted successfully"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRail::postSend(uint64_t immediate_data, + fi_addr_t dest_addr, + nixlLibfabricReq *req) const { + if (req->buffer_size == 0 || req->buffer_size > ControlRequestPool::BUFFER_SIZE) { + NIXL_ERROR << "Invalid message size: " << req->buffer_size + << " (max: " << ControlRequestPool::BUFFER_SIZE << ")"; + return NIXL_ERR_INVALID_PARAM; + } + + // Prepare descriptor + void *desc = fi_mr_desc(req->mr); + + NIXL_TRACE << "Sending data on endpoint: " << endpoint << " buffer: " << req->buffer + << " size: " << req->buffer_size << " immediate_data: " << std::hex << immediate_data + << " msg_type: " << NIXL_GET_MSG_TYPE_FROM_IMM(immediate_data) + << " agent_idx: " << NIXL_GET_AGENT_INDEX_FROM_IMM(immediate_data) + << " XFER_ID: " << NIXL_GET_XFER_ID_FROM_IMM(immediate_data) + << " dest_addr: " << dest_addr << std::dec << " context: " << &req->ctx; + + // Libfabric fi_senddata call + int ret = fi_senddata( + endpoint, req->buffer, req->buffer_size, desc, immediate_data, dest_addr, &req->ctx); + if (ret) { + NIXL_ERROR << "fi_senddata failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + NIXL_TRACE << "Send posted successfully"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRail::postWrite(const void *local_buffer, + size_t length, + void *local_desc, + uint64_t immediate_data, + fi_addr_t dest_addr, + uint64_t remote_addr, + uint64_t remote_key, + nixlLibfabricReq *req) const { + // Validation + if (!req) { + NIXL_ERROR << "Invalid request for write on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + + // Setup and logging + NIXL_TRACE << "Posting RDMA write on endpoint: " << std::hex << endpoint + << " local_buffer: " << local_buffer << " length: " << length + << " immediate_data: " << immediate_data << " dest_addr: " << dest_addr + << " remote_addr: " << (void *)remote_addr << " remote_key: " << remote_key + << " context: " << &req->ctx; + + // Libfabric fi_writedata call + int ret = fi_writedata(endpoint, + local_buffer, + length, + local_desc, + immediate_data, + dest_addr, + remote_addr, + remote_key, + &req->ctx); + + if (ret) { + NIXL_ERROR << "fi_writedata failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + NIXL_TRACE << "RDMA write posted successfully"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRail::postRead(void *local_buffer, + size_t length, + void *local_desc, + fi_addr_t dest_addr, + uint64_t remote_addr, + uint64_t remote_key, + nixlLibfabricReq *req) const { + if (!req) { + NIXL_ERROR << "Invalid request for read on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_TRACE << "Posting RDMA read on endpoint: " << std::hex << endpoint + << " local_buffer: " << local_buffer << " length: " << length + << " dest_addr: " << dest_addr << " remote_addr: " << (void *)remote_addr + << " remote_key: " << remote_key << " context: " << &req->ctx; + + int ret = fi_read( + endpoint, local_buffer, length, local_desc, dest_addr, remote_addr, remote_key, &req->ctx); + if (ret) { + NIXL_ERROR << "fi_read failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + NIXL_TRACE << "RDMA read posted successfully"; + return NIXL_SUCCESS; +} + +// Memory Registration Methods + +nixl_status_t +nixlLibfabricRail::registerMemory(void *buffer, + size_t length, + uint64_t access_flags, + struct fid_mr **mr_out, + uint64_t *key_out) const { + if (!buffer || !mr_out || !key_out) { + NIXL_ERROR << "Invalid parameters on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + if (!domain) { + NIXL_ERROR << "Domain not initialized on rail " << rail_id; + return NIXL_ERR_BACKEND; + } + + struct fid_mr *mr; + int ret = fi_mr_reg(domain, buffer, length, access_flags, 0, 0, 0, &mr, NULL); + if (ret) { + NIXL_ERROR << "fi_mr_reg failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + + *mr_out = mr; + *key_out = fi_mr_key(mr); + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRail::deregisterMemory(struct fid_mr *mr) const { + if (!mr) { + NIXL_ERROR << "Invalid MR parameter on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + + int ret = fi_close(&mr->fid); + if (ret) { + NIXL_ERROR << "fi_close failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + + return NIXL_SUCCESS; +} + +// Address Vector Management Methods + +nixl_status_t +nixlLibfabricRail::insertAddress(const void *addr, fi_addr_t *fi_addr_out) const { + if (!addr || !fi_addr_out) { + NIXL_ERROR << "Invalid parameters on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + if (!av) { + NIXL_ERROR << "Address vector not initialized on rail " << rail_id; + return NIXL_ERR_BACKEND; + } + + int ret = fi_av_insert(av, addr, 1, fi_addr_out, 0, NULL); + if (ret != 1) { + NIXL_ERROR << "fi_av_insert failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRail::removeAddress(fi_addr_t fi_addr) const { + if (fi_addr == FI_ADDR_UNSPEC) { + NIXL_ERROR << "Invalid fi_addr parameter on rail " << rail_id; + return NIXL_ERR_INVALID_PARAM; + } + if (!av) { + NIXL_ERROR << "Address vector not initialized on rail " << rail_id; + return NIXL_ERR_BACKEND; + } + + int ret = fi_av_remove(av, &fi_addr, 1, 0); + if (ret != 0) { + NIXL_ERROR << "fi_av_remove failed on rail " << rail_id << ": " << fi_strerror(-ret); + return NIXL_ERR_BACKEND; + } + + return NIXL_SUCCESS; +} + +// Memory Descriptor Helper Methods + +void * +nixlLibfabricRail::getMemoryDescriptor(struct fid_mr *mr) const { + if (!mr) { + NIXL_ERROR << "Invalid MR parameter on rail " << rail_id; + return nullptr; + } + return fi_mr_desc(mr); +} + +uint64_t +nixlLibfabricRail::getMemoryKey(struct fid_mr *mr) const { + if (!mr) { + NIXL_ERROR << "Invalid MR parameter on rail " << rail_id; + return 0; + } + return fi_mr_key(mr); +} + +// Optimized Resource Management Methods + +nixlLibfabricReq * +nixlLibfabricRail::allocateControlRequest(size_t needed_size) { + return control_request_pool_.allocate(needed_size); +} + +nixlLibfabricReq * +nixlLibfabricRail::allocateDataRequest(nixlLibfabricReq::OpType op_type) { + return data_request_pool_.allocate(op_type); +} + +void +nixlLibfabricRail::releaseRequest(nixlLibfabricReq *req) { + if (!req) { + NIXL_ERROR << "Null request provided to releaseRequest on rail " << rail_id; + return; + } + // Determine which pool to release to based on operation type + if (req->operation_type == nixlLibfabricReq::SEND || + req->operation_type == nixlLibfabricReq::RECV) { + control_request_pool_.release(req); + } else { + data_request_pool_.release(req); + } + NIXL_TRACE << "Released request with XFER_ID " << req->xfer_id; +} + +nixlLibfabricReq * +nixlLibfabricRail::findRequestFromContext(void *context) const { + if (!context) { + NIXL_ERROR << "Null context provided to findRequestFromContext on rail " << rail_id; + return nullptr; + } + // Try control pool first + nixlLibfabricReq *req = control_request_pool_.findByContext(context); + if (req) { + return req; + } + // Try data pool + req = data_request_pool_.findByContext(context); + if (req) { + return req; + } + NIXL_ERROR << "No request found for context " << context << " on rail " << rail_id; + return nullptr; +} diff --git a/src/utils/libfabric/libfabric_rail.h b/src/utils/libfabric/libfabric_rail.h new file mode 100644 index 0000000000..781f7ca99a --- /dev/null +++ b/src/utils/libfabric/libfabric_rail.h @@ -0,0 +1,381 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_H +#define NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_H + +#include +#include +#include +#include +#include +#include + +#include "nixl.h" +#include "backend/backend_aux.h" +#include "libfabric/libfabric_common.h" + +// Forward declarations +class nixlLibfabricConnection; + +/** + * @brief Request structure for libfabric operations + * + */ +struct nixlLibfabricReq { + fi_context2 ctx; ///< Libfabric context for operation tracking + size_t rail_id; ///< Rail ID that owns this request + uint32_t xfer_id; ///< Pre-assigned globally unique transfer ID + void *buffer; ///< Pre-assigned buffer for CONTROL operations, nullptr for DATA + struct fid_mr *mr; ///< Pre-assigned memory registration for CONTROL, nullptr for DATA + size_t buffer_size; ///< Pre-assigned buffer size for CONTROL (2KB), 0 for DATA + + enum OpType { WRITE, READ, SEND, RECV } operation_type; ///< Operation type (pre-assigned) + + bool in_use; ///< Pool management flag + size_t chunk_offset; ///< Chunk offset for DATA requests + size_t chunk_size; ///< Chunk size for DATA requests + std::function completion_callback; ///< Completion callback function + void *local_addr; ///< Local memory address for transfers + uint64_t remote_addr; ///< Remote memory address for transfers + struct fid_mr *local_mr; ///< Local memory registration for transfers + uint64_t remote_key; ///< Remote access key for transfers + + /** Default constructor initializing all fields */ + nixlLibfabricReq() + : rail_id(0), + xfer_id(0), + buffer(nullptr), + mr(nullptr), + buffer_size(0), + operation_type(SEND), + in_use(false), + chunk_offset(0), + chunk_size(0), + local_addr(nullptr), + remote_addr(0), + local_mr(nullptr), + remote_key(0) { + memset(&ctx, 0, sizeof(fi_context2)); + } +}; + +/** Thread-safe request pool with O(1) allocation/release */ +class RequestPool { +public: + /** Initialize request pool with specified size */ + RequestPool(size_t pool_size, size_t rail_id); + + /** Virtual destructor for proper cleanup */ + virtual ~RequestPool() = default; + + /** Release request back to the pool */ + virtual void + release(nixlLibfabricReq *req); + + /** Find request by libfabric context pointer */ + nixlLibfabricReq * + findByContext(void *context) const; + + /** Get count of currently active requests */ + size_t + getActiveRequestCount() const; + + // Non-copyable and non-movable since we use unique_ptr for management + RequestPool(const RequestPool &) = delete; + RequestPool & + operator=(const RequestPool &) = delete; + RequestPool(RequestPool &&) = delete; + RequestPool & + operator=(RequestPool &&) = delete; + +protected: + std::vector requests_; ///< Fixed-size request pool + std::stack free_indices_; ///< Stack of available request indices + size_t rail_id_; ///< Rail ID for this pool + mutable std::mutex pool_mutex_; ///< Thread safety protection +}; + +/** Control request pool with pre-allocated buffers for SEND/RECV operations */ +class ControlRequestPool : public RequestPool { +public: + /** Initialize control request pool */ + ControlRequestPool(size_t pool_size, size_t rail_id); + + /** Destructor with explicit cleanup */ + ~ControlRequestPool(); + + // Non-copyable and non-movable since we use unique_ptr for management + ControlRequestPool(const ControlRequestPool &) = delete; + ControlRequestPool & + operator=(const ControlRequestPool &) = delete; + ControlRequestPool(ControlRequestPool &&) = delete; + ControlRequestPool & + operator=(ControlRequestPool &&) = delete; + + /** Initialize pool with buffers and pre-assigned XFER_IDs */ + nixl_status_t + initializeWithBuffersAndXferIds(struct fid_domain *domain, + const std::vector &xfer_ids); + + /** Allocate control request with size validation */ + nixlLibfabricReq * + allocate(size_t needed_size); + + /** Explicit cleanup method for proper resource ordering */ + void + cleanup(); + + /** Buffer size constant for validation */ + static constexpr size_t BUFFER_SIZE = NIXL_LIBFABRIC_SEND_RECV_BUFFER_SIZE; + +private: + void *buffer_chunk_; ///< Large pre-registered buffer chunk + size_t buffer_chunk_size_; ///< Total size of buffer chunk + struct fid_mr *buffer_mr_; ///< Memory registration for chunk +}; + +/** Lightweight data request pool for WRITE/READ operations */ +class DataRequestPool : public RequestPool { +public: + /** Initialize data request pool */ + DataRequestPool(size_t pool_size, size_t rail_id); + + /** Default destructor (no special cleanup needed) */ + ~DataRequestPool() = default; + + // Non-copyable and non-movable since we use unique_ptr for management + DataRequestPool(const DataRequestPool &) = delete; + DataRequestPool & + operator=(const DataRequestPool &) = delete; + DataRequestPool(DataRequestPool &&) = delete; + DataRequestPool & + operator=(DataRequestPool &&) = delete; + + /** Initialize pool with pre-assigned XFER_IDs */ + nixl_status_t + initializeWithXferIds(const std::vector &xfer_ids); + + /** Allocate data request for specified operation type */ + nixlLibfabricReq * + allocate(nixlLibfabricReq::OpType op_type); +}; + + +/** Connection state tracking for multi-rail connections */ +enum class ConnectionState { + DISCONNECTED, ///< No connection attempt made, initial state + CONNECT_REQ_SENT, ///< Connection request sent, waiting for ACK + CONNECT_ACK_SENT, ///< Connection ACK sent (target side) + CONNECTED, ///< ACK received, ready for data transfers + FAILED ///< Connection attempt failed +}; + +// Stream operator for ConnectionState to enable logging +inline std::ostream & +operator<<(std::ostream &os, const ConnectionState &state) { + switch (state) { + case ConnectionState::DISCONNECTED: + return os << "DISCONNECTED"; + case ConnectionState::CONNECT_REQ_SENT: + return os << "CONNECT_REQ_SENT"; + case ConnectionState::CONNECT_ACK_SENT: + return os << "CONNECT_ACK_SENT"; + case ConnectionState::CONNECTED: + return os << "CONNECTED"; + case ConnectionState::FAILED: + return os << "FAILED"; + default: + return os << "UNKNOWN"; + } +} + +/** Individual libfabric rail managing fabric, domain, endpoint, CQ, and AV */ +class nixlLibfabricRail { +public: + uint16_t rail_id; ///< Unique rail identifier + std::string device_name; ///< EFA device name for this rail + char ep_name[LF_EP_NAME_MAX_LEN]; ///< Endpoint name for connection setup + mutable bool blocking_cq_sread_supported; ///< Whether blocking CQ reads are supported + struct fid_ep *endpoint; ///< Libfabric endpoint handle + + /** Initialize libfabric rail with all resources */ + nixlLibfabricRail(const std::string &device, const std::string &provider, uint16_t id); + + /** Destroy rail and cleanup all libfabric resources */ + ~nixlLibfabricRail(); + + // Non-copyable and non-movable since we use unique_ptr for management + nixlLibfabricRail(const nixlLibfabricRail &) = delete; + nixlLibfabricRail & + operator=(const nixlLibfabricRail &) = delete; + nixlLibfabricRail(nixlLibfabricRail &&) = delete; + nixlLibfabricRail & + operator=(nixlLibfabricRail &&) = delete; + + /** Explicit cleanup method for proper resource ordering */ + void + cleanup(); + + /** Validate that rail is properly initialized */ + bool + isProperlyInitialized() const; + + // Memory registration methods + /** Register memory buffer with libfabric */ + nixl_status_t + registerMemory(void *buffer, + size_t length, + uint64_t access_flags, + struct fid_mr **mr_out, + uint64_t *key_out) const; + + /** Deregister memory from libfabric */ + nixl_status_t + deregisterMemory(struct fid_mr *mr) const; + + // Address vector management methods + /** Insert remote endpoint address into address vector */ + nixl_status_t + insertAddress(const void *addr, fi_addr_t *fi_addr_out) const; + + /** Remove address from address vector */ + nixl_status_t + removeAddress(fi_addr_t fi_addr) const; + + // Memory descriptor helper methods + /** Get libfabric memory descriptor for MR */ + void * + getMemoryDescriptor(struct fid_mr *mr) const; + + /** Get remote access key for MR */ + uint64_t + getMemoryKey(struct fid_mr *mr) const; + + // Libfabric operation wrappers + /** Post receive operation */ + nixl_status_t + postRecv(nixlLibfabricReq *req) const; + + /** Post send operation with immediate data */ + nixl_status_t + postSend(uint64_t immediate_data, fi_addr_t dest_addr, nixlLibfabricReq *req) const; + + /** Post RDMA write operation with immediate data */ + nixl_status_t + postWrite(const void *local_buffer, + size_t length, + void *local_desc, + uint64_t immediate_data, + fi_addr_t dest_addr, + uint64_t remote_addr, + uint64_t remote_key, + nixlLibfabricReq *req) const; + + /** Post RDMA read operation */ + nixl_status_t + postRead(void *local_buffer, + size_t length, + void *local_desc, + fi_addr_t dest_addr, + uint64_t remote_addr, + uint64_t remote_key, + nixlLibfabricReq *req) const; + + /** Process completion queue with batching support */ + nixl_status_t + progressCompletionQueue(bool use_blocking = false); + + // Callback registration methods + /** Set callback for notification message processing */ + void + setNotificationCallback(std::function callback); + + /** Set callback for connection acknowledgment processing */ + void + setConnectionAckCallback( + std::function callback); + + /** Set callback for connection request processing */ + void + setConnectionReqCallback( + std::function callback); + + /** Set callback for XFER_ID tracking */ + void + setXferIdCallback(std::function callback); + + // Optimized resource management methods + /** Allocate control request with size validation */ + [[nodiscard]] nixlLibfabricReq * + allocateControlRequest(size_t needed_size); + + /** Allocate data request for specified operation */ + [[nodiscard]] nixlLibfabricReq * + allocateDataRequest(nixlLibfabricReq::OpType op_type); + + /** Release request back to appropriate pool */ + void + releaseRequest(nixlLibfabricReq *req); + + /** Find request from libfabric context pointer */ + nixlLibfabricReq * + findRequestFromContext(void *context) const; + +private: + // Core libfabric resources + struct fi_info *info; // from rail_infos[rail_id] + struct fid_fabric *fabric; // from rail_fabrics[rail_id] + struct fid_domain *domain; // from rail_domains[rail_id] + struct fid_cq *cq; // from rail_cqs[rail_id] + struct fid_av *av; // from rail_avs[rail_id] + + // CQ progress mutex to protect completion queue operations + mutable std::mutex cq_progress_mutex_; + + // Callback functions + std::function notificationCallback; + std::function connectionAckCallback; + std::function + connectionReqCallback; + // XFER_ID tracking callback + std::function xferIdCallback; + + // Separate request pools for optimal performance + ControlRequestPool + control_request_pool_; // 256 CONTROL requests (SEND/RECV) with internal buffers + DataRequestPool data_request_pool_; // 1024 DATA requests (WRITE/read) - no buffers needed + + // Configuration constants + static constexpr size_t CONTROL_REQUESTS_PER_RAIL = + 256; // SEND/RECV operations (1:1 with buffers) + static constexpr size_t DATA_REQUESTS_PER_RAIL = 1024; // WRITE/read operations (no buffers) + + nixl_status_t + processCompletionQueueEntry(struct fi_cq_data_entry *comp); + nixl_status_t + processLocalSendCompletion(struct fi_cq_data_entry *comp); + nixl_status_t + processLocalTransferCompletion(struct fi_cq_data_entry *comp, const char *operation_type); + nixl_status_t + processRecvCompletion(struct fi_cq_data_entry *comp); + nixl_status_t + processRemoteWriteCompletion(struct fi_cq_data_entry *comp) const; +}; + + +#endif // NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_H diff --git a/src/utils/libfabric/libfabric_rail_manager.cpp b/src/utils/libfabric/libfabric_rail_manager.cpp new file mode 100644 index 0000000000..6db3495471 --- /dev/null +++ b/src/utils/libfabric/libfabric_rail_manager.cpp @@ -0,0 +1,855 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric_rail_manager.h" +#include "libfabric/libfabric_common.h" +#include "libfabric/libfabric_topology.h" +#include "common/nixl_log.h" +#include "serdes/serdes.h" + +// Static round-robin counter for rail selection +static std::atomic round_robin_counter{0}; + +nixlLibfabricRailManager::nixlLibfabricRailManager(size_t striping_threshold) + : striping_threshold_(striping_threshold) { + NIXL_DEBUG << "Creating rail manager with striping threshold: " << striping_threshold_ + << " bytes"; + + // Initialize topology system + try { + topology = std::make_unique(); + NIXL_DEBUG << "System topology discovered successfully"; + } + catch (const std::exception &e) { + NIXL_ERROR << "Failed to discover system topology: " << e.what(); + throw std::runtime_error( + "Topology discovery failed - cannot proceed without topology information"); + } + + // Get EFA devices from topology and create rails automatically + std::vector all_efa_devices = topology->getAllEfaDevices(); + std::string selected_fabric_name = topology->getEFAfabricName(); + + NIXL_DEBUG << "Got " << all_efa_devices.size() << " EFA devices from topology for the fabric" + << selected_fabric_name; + + // Create data rails with selected provider + nixl_status_t rail_status = createDataRails(all_efa_devices, selected_fabric_name); + if (rail_status != NIXL_SUCCESS) { + throw std::runtime_error("Rail Manager failed to create data rails"); + } + // Create control rails with selected provider + nixl_status_t control_rail_status = createControlRails( + all_efa_devices, selected_fabric_name, NIXL_LIBFABRIC_DEFAULT_CONTROL_RAILS); + if (control_rail_status != NIXL_SUCCESS) { + throw std::runtime_error("Rail Manager failed to create control rails"); + } + NIXL_DEBUG << "Successfully created " << data_rails_.size() << " data rails and " + << control_rails_.size() << " control rails"; +} + +nixlLibfabricRailManager::~nixlLibfabricRailManager() { + NIXL_DEBUG << "Destroying rail manager"; +} + +nixl_status_t +nixlLibfabricRailManager::createDataRails(const std::vector &efa_devices, + const std::string &provider_name) { + num_data_rails_ = efa_devices.size(); + // Pre-allocate to ensure contiguous memory allocation + data_rails_.reserve(num_data_rails_); + + // Build EFA device to rail index mapping for O(1) lookup + efa_device_to_rail_map.reserve(num_data_rails_); + + try { + data_rails_.clear(); + data_rails_.reserve(num_data_rails_); + + for (size_t i = 0; i < num_data_rails_; ++i) { + data_rails_.emplace_back(std::make_unique( + efa_devices[i], provider_name, static_cast(i))); + + // Initialize EFA device mapping + efa_device_to_rail_map[efa_devices[i]] = i; + + NIXL_DEBUG << "Created data rail " << i << " (device: " << efa_devices[i] + << ", provider: " << provider_name << ")"; + } + } + catch (const std::exception &e) { + NIXL_ERROR << "Failed to create data rails: " << e.what(); + return NIXL_ERR_BACKEND; + } + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::createControlRails(const std::vector &efa_devices, + const std::string &provider_name, + size_t num_control_rails) { + // Pre-allocate to ensure contiguous memory allocation + num_control_rails_ = num_control_rails; + control_rails_.reserve(num_control_rails_); + + try { + control_rails_.clear(); + control_rails_.reserve(num_control_rails_); + + for (size_t i = 0; i < num_control_rails_; ++i) { + control_rails_.emplace_back(std::make_unique( + efa_devices[i], provider_name, static_cast(i))); + NIXL_DEBUG << "Created control rail " << i << " (device: " << efa_devices[i] + << ", provider: " << provider_name << ")"; + } + } + catch (const std::exception &e) { + NIXL_ERROR << "Failed to create control rails: " << e.what(); + return NIXL_ERR_BACKEND; + } + return NIXL_SUCCESS; +} + +bool +nixlLibfabricRailManager::shouldUseStriping(size_t transfer_size) const { + return transfer_size >= striping_threshold_; +} + +nixl_status_t +nixlLibfabricRailManager::prepareAndSubmitTransfer(nixlLibfabricReq::OpType op_type, + void *local_addr, + size_t transfer_size, + uint64_t remote_base_addr, + const std::vector &selected_rails, + const std::vector &local_mrs, + const std::vector &remote_keys, + const std::vector &dest_addrs, + uint16_t agent_idx, + std::function completion_callback, + BinaryNotification *binary_notif) { + if (selected_rails.empty()) { + NIXL_ERROR << "No rails selected for transfer"; + return NIXL_ERR_INVALID_PARAM; + } + + // Determine striping strategy + bool use_striping = shouldUseStriping(transfer_size) && selected_rails.size() > 1; + + if (!use_striping) { + // Round-robin: use one rail for entire transfer + size_t rail_idx = round_robin_counter.fetch_add(1) % selected_rails.size(); + size_t rail_id = selected_rails[rail_idx]; + // Allocate request + nixlLibfabricReq *req = data_rails_[rail_id]->allocateDataRequest(op_type); + if (!req) { + NIXL_ERROR << "Failed to allocate request for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + // Set completion callback and populate request + req->completion_callback = completion_callback; + req->chunk_offset = 0; + req->chunk_size = transfer_size; + req->local_addr = local_addr; + req->remote_addr = remote_base_addr; + req->local_mr = local_mrs[rail_id]; + req->remote_key = remote_keys[rail_id]; + req->rail_id = rail_id; + // Submit immediately + nixl_status_t status; + if (op_type == nixlLibfabricReq::WRITE) { + uint64_t imm_data = + NIXL_MAKE_IMM_DATA(NIXL_LIBFABRIC_MSG_TRANSFER, agent_idx, req->xfer_id); + status = data_rails_[rail_id]->postWrite(req->local_addr, + req->chunk_size, + fi_mr_desc(req->local_mr), + imm_data, + dest_addrs[rail_id], + req->remote_addr, + req->remote_key, + req); + } else { + status = data_rails_[rail_id]->postRead(req->local_addr, + req->chunk_size, + fi_mr_desc(req->local_mr), + dest_addrs[rail_id], + req->remote_addr, + req->remote_key, + req); + } + if (status != NIXL_SUCCESS) { + data_rails_[rail_id]->releaseRequest(req); + return status; + } + + // Collect XFER_ID directly in BinaryNotification + if (binary_notif && binary_notif->xfer_id_count < NIXL_LIBFABRIC_MAX_XFER_IDS) { + binary_notif->addXferId(req->xfer_id); + } + + NIXL_DEBUG << "Round-robin: submitted single request on rail " << rail_id << " for " + << transfer_size << " bytes, XFER_ID: " << req->xfer_id; + + } else { + // Striping: distribute across multiple rails + size_t num_rails = selected_rails.size(); + size_t chunk_size = transfer_size / num_rails; + size_t remainder = transfer_size % num_rails; + for (size_t i = 0; i < num_rails; ++i) { + size_t rail_id = selected_rails[i]; + size_t current_chunk_size = chunk_size + (i == num_rails - 1 ? remainder : 0); + if (current_chunk_size == 0) break; + // Allocate request + nixlLibfabricReq *req = data_rails_[rail_id]->allocateDataRequest(op_type); + if (!req) { + NIXL_ERROR << "Failed to allocate request for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + + req->completion_callback = completion_callback; + + // Calculate and populate chunk info + size_t chunk_offset = i * chunk_size; + req->chunk_offset = chunk_offset; + req->chunk_size = current_chunk_size; + req->local_addr = static_cast(local_addr) + chunk_offset; + req->remote_addr = remote_base_addr + chunk_offset; + req->local_mr = local_mrs[rail_id]; + req->remote_key = remote_keys[rail_id]; + req->rail_id = rail_id; + nixl_status_t status; + if (op_type == nixlLibfabricReq::WRITE) { + uint64_t imm_data = + NIXL_MAKE_IMM_DATA(NIXL_LIBFABRIC_MSG_TRANSFER, agent_idx, req->xfer_id); + status = data_rails_[rail_id]->postWrite(req->local_addr, + req->chunk_size, + fi_mr_desc(req->local_mr), + imm_data, + dest_addrs[rail_id], + req->remote_addr, + req->remote_key, + req); + } else { + status = data_rails_[rail_id]->postRead(req->local_addr, + req->chunk_size, + fi_mr_desc(req->local_mr), + dest_addrs[rail_id], + req->remote_addr, + req->remote_key, + req); + } + if (status != NIXL_SUCCESS) { + data_rails_[rail_id]->releaseRequest(req); + return status; + } + + // Collect XFER_ID directly in BinaryNotification + if (binary_notif && binary_notif->xfer_id_count < NIXL_LIBFABRIC_MAX_XFER_IDS) { + binary_notif->addXferId(req->xfer_id); + } + } + NIXL_DEBUG << "Striping: submitted " << (binary_notif ? binary_notif->xfer_id_count : 0) + << " requests for " << transfer_size << " bytes"; + } + + NIXL_DEBUG << "Successfully submitted " << (binary_notif ? binary_notif->xfer_id_count : 0) + << " requests for " << transfer_size << " bytes"; + + return NIXL_SUCCESS; +} + +std::vector +nixlLibfabricRailManager::selectRailsForMemory(void *mem_addr, + nixl_mem_t mem_type, + int gpu_id) const { + if (mem_type == VRAM_SEG) { +#ifdef HAVE_CUDA + if (gpu_id < 0) { + NIXL_ERROR << "Invalid GPU ID " << gpu_id << " for VRAM memory " << mem_addr; + return {}; // Return empty vector to indicate failure + } + std::vector gpu_efa_devices = topology->getEfaDevicesForGpu(gpu_id); + if (gpu_efa_devices.empty()) { + NIXL_ERROR << "No EFA devices found for GPU " << gpu_id; + return {}; // Return empty vector to indicate failure + } + std::vector gpu_rails; + for (const std::string &efa_device : gpu_efa_devices) { + auto it = efa_device_to_rail_map.find(efa_device); + if (it != efa_device_to_rail_map.end()) { + // Bounds check: ensure rail index is valid + if (it->second < data_rails_.size()) { + gpu_rails.push_back(it->second); + NIXL_DEBUG << "VRAM memory " << mem_addr << " on GPU " << gpu_id + << " mapped to rail " << it->second << " (EFA device: " << efa_device + << ")"; + } else { + NIXL_WARN << "EFA device " << efa_device << " maps to rail " << it->second + << " but only " << data_rails_.size() << " rails available"; + } + } else { + NIXL_WARN << "EFA device " << efa_device << " not found in rail mapping for GPU " + << gpu_id; + } + } + + if (gpu_rails.empty()) { + NIXL_ERROR << "No valid rail mapping found for GPU " << gpu_id << " (checked " + << gpu_efa_devices.size() << " EFA devices)"; + return {}; + } + + NIXL_DEBUG << "VRAM memory " << mem_addr << " on GPU " << gpu_id << " will use " + << gpu_rails.size() << " rails total"; + return gpu_rails; +#else + NIXL_ERROR << "VRAM memory type not supported without CUDA"; + return {}; +#endif + } + if (mem_type == DRAM_SEG) { + // For DRAM, use all available rails for maximum bandwidth + std::vector all_rails; + all_rails.reserve(data_rails_.size()); + for (size_t i = 0; i < data_rails_.size(); ++i) { + all_rails.push_back(i); + } + + NIXL_DEBUG << "DRAM memory " << mem_addr << " will use all " << all_rails.size() + << " available rails for maximum bandwidth"; + return all_rails; + } + + // For unsupported memory types, return empty vector + NIXL_ERROR << "Unsupported memory type " << mem_type; + return {}; +} + +nixl_status_t +nixlLibfabricRailManager::registerMemory(void *buffer, + size_t length, + nixl_mem_t mem_type, + int gpu_id, + std::vector &mr_list_out, + std::vector &key_list_out, + std::vector &selected_rails_out) { + if (!buffer) { + NIXL_ERROR << "Invalid buffer parameter"; + return NIXL_ERR_INVALID_PARAM; + } + + // Use internal rail selection with explicit GPU ID + std::vector selected_rails = selectRailsForMemory(buffer, mem_type, gpu_id); + if (selected_rails.empty()) { + NIXL_ERROR << "No rails selected for memory type " << mem_type; + return NIXL_ERR_NOT_SUPPORTED; + } + + // Resize output vectors to match all rails + mr_list_out.resize(data_rails_.size(), nullptr); + key_list_out.resize(data_rails_.size(), 0); + selected_rails_out = selected_rails; // Return which rails were selected + + // Register memory on each selected rail + for (size_t i = 0; i < selected_rails.size(); ++i) { + size_t rail_idx = selected_rails[i]; + if (rail_idx >= data_rails_.size()) { + NIXL_ERROR << "Invalid rail index " << rail_idx; + // Cleanup already registered MRs + for (size_t cleanup_idx : selected_rails) { + if (cleanup_idx >= rail_idx) break; // Only cleanup what we've done so far + if (mr_list_out[cleanup_idx]) { + data_rails_[cleanup_idx]->deregisterMemory(mr_list_out[cleanup_idx]); + mr_list_out[cleanup_idx] = nullptr; + } + } + return NIXL_ERR_INVALID_PARAM; + } + + struct fid_mr *mr; + uint64_t key; + nixl_status_t status = data_rails_[rail_idx]->registerMemory( + buffer, length, FI_REMOTE_WRITE | FI_REMOTE_READ, &mr, &key); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to register memory on rail " << rail_idx; + // Cleanup already registered MRs + for (size_t cleanup_idx : selected_rails) { + if (cleanup_idx >= rail_idx) break; // Only cleanup what we've done so far + if (mr_list_out[cleanup_idx]) { + data_rails_[cleanup_idx]->deregisterMemory(mr_list_out[cleanup_idx]); + mr_list_out[cleanup_idx] = nullptr; + } + } + return status; + } + + mr_list_out[rail_idx] = mr; + key_list_out[rail_idx] = key; + + // Mark rail as active for progress tracking optimization + markRailActive(rail_idx); + + NIXL_DEBUG << "Registered memory on rail " << rail_idx + << " (mr: " << reinterpret_cast(mr) << ", key: " << key << ")"; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::deregisterMemory(const std::vector &selected_rails, + const std::vector &mr_list) { + if (selected_rails.empty() || mr_list.size() != data_rails_.size()) { + NIXL_ERROR << "Invalid parameters"; + return NIXL_ERR_INVALID_PARAM; + } + + nixl_status_t overall_status = NIXL_SUCCESS; + + for (size_t i = 0; i < selected_rails.size(); ++i) { + size_t rail_idx = selected_rails[i]; + if (rail_idx >= data_rails_.size()) { + NIXL_ERROR << "Invalid rail index " << rail_idx; + overall_status = NIXL_ERR_INVALID_PARAM; + continue; + } + + if (mr_list[rail_idx]) { + nixl_status_t status = data_rails_[rail_idx]->deregisterMemory(mr_list[rail_idx]); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to deregister memory on rail " << rail_idx; + overall_status = status; + } + markRailInactive(rail_idx); + } + } + + return overall_status; +} + +nixl_status_t +nixlLibfabricRailManager::insertAllAddresses( + RailType rail_type, + const std::vector> &endpoints, + std::vector &fi_addrs_out, + std::vector &ep_names_out) { + auto &rails = (rail_type == RailType::DATA) ? data_rails_ : control_rails_; + const char *rail_type_str = (rail_type == RailType::DATA) ? "data" : "control"; + + if (endpoints.size() != rails.size()) { + NIXL_ERROR << "Expected " << rails.size() << " " << rail_type_str << " endpoints, got " + << endpoints.size(); + return NIXL_ERR_INVALID_PARAM; + } + + fi_addrs_out.clear(); + ep_names_out.clear(); + fi_addrs_out.reserve(rails.size()); + ep_names_out.reserve(rails.size()); + + // Process all rails in one operation + for (size_t rail_id = 0; rail_id < rails.size(); ++rail_id) { + fi_addr_t fi_addr; + nixl_status_t status = rails[rail_id]->insertAddress(endpoints[rail_id].data(), &fi_addr); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed for " << rail_type_str << " rail " << rail_id; + return status; + } + + fi_addrs_out.push_back(fi_addr); + ep_names_out.push_back( + rails[rail_id] + ->ep_name); // This is char[LF_EP_NAME_MAX_LEN], will be converted to char* + + NIXL_DEBUG << "Processed " << rail_type_str << " rail " << rail_id + << " (fi_addr: " << fi_addr << ")"; + } + + NIXL_DEBUG << "Successfully processed " << rails.size() << " " << rail_type_str << " rails"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::cleanupConnection(RailType rail_type, + const std::vector &fi_addrs_to_remove) { + auto &rails = (rail_type == RailType::DATA) ? data_rails_ : control_rails_; + const char *rail_type_str = (rail_type == RailType::DATA) ? "data" : "control"; + + if (fi_addrs_to_remove.size() != rails.size()) { + NIXL_ERROR << "Expected " << rails.size() << " " << rail_type_str << " fi_addrs, got " + << fi_addrs_to_remove.size(); + return NIXL_ERR_INVALID_PARAM; + } + + NIXL_DEBUG << "Cleaning up connection for " << rails.size() << " " << rail_type_str << " rails"; + // Remove addresses from all rails + nixl_status_t overall_status = NIXL_SUCCESS; + for (size_t rail_id = 0; rail_id < rails.size(); ++rail_id) { + if (fi_addrs_to_remove[rail_id] != FI_ADDR_UNSPEC) { + nixl_status_t status = rails[rail_id]->removeAddress(fi_addrs_to_remove[rail_id]); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to remove address from " << rail_type_str << " rail " + << rail_id << ", fi_addr: " << fi_addrs_to_remove[rail_id]; + overall_status = status; + // Continue cleanup for other rails even if one fails + } else { + NIXL_DEBUG << "Successfully removed address from " << rail_type_str << " rail " + << rail_id << ", fi_addr: " << fi_addrs_to_remove[rail_id]; + } + } else { + NIXL_DEBUG << "Skipping FI_ADDR_UNSPEC for " << rail_type_str << " rail " << rail_id; + } + } + NIXL_DEBUG << "Completed cleanup for " << rails.size() << " " << rail_type_str << " rails"; + return overall_status; +} + +nixl_status_t +nixlLibfabricRailManager::postControlMessage(ControlMessageType msg_type, + nixlLibfabricReq *req, + fi_addr_t dest_addr, + uint16_t agent_idx, + std::function completion_callback) { + // Validation + if (control_rails_.empty()) { + NIXL_ERROR << "No control rails available"; + return NIXL_ERR_INVALID_PARAM; + } + + if (!req) { + NIXL_ERROR << "Pre-allocated request is null"; + return NIXL_ERR_INVALID_PARAM; + } + + uint64_t msg_type_value; + switch (msg_type) { + case ControlMessageType::NOTIFICATION: + msg_type_value = NIXL_LIBFABRIC_MSG_NOTIFICTION; + break; + case ControlMessageType::CONNECTION_REQ: + msg_type_value = NIXL_LIBFABRIC_MSG_CONNECT; + break; + case ControlMessageType::CONNECTION_ACK: + msg_type_value = NIXL_LIBFABRIC_MSG_ACK; + break; + case ControlMessageType::DISCONNECT_REQ: + msg_type_value = NIXL_LIBFABRIC_MSG_DISCONNECT; + break; + default: + NIXL_ERROR << "Unknown message type"; + return NIXL_ERR_INVALID_PARAM; + } + size_t control_rail_id = 0; + uint32_t xfer_id = req->xfer_id; + uint64_t imm_data = NIXL_MAKE_IMM_DATA(msg_type_value, agent_idx, xfer_id); + + // Set completion callback if provided + if (completion_callback) { + req->completion_callback = completion_callback; + NIXL_DEBUG << "Set completion callback for control message request " << req->xfer_id; + } + + NIXL_DEBUG << "Sending control message type " << msg_type_value << " agent_idx=" << agent_idx + << " XFER_ID=" << xfer_id << " imm_data=0x" << std::hex << imm_data << std::dec; + + // Rail postSend + nixl_status_t status = control_rails_[control_rail_id]->postSend(imm_data, dest_addr, req); + + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to send control message type " << static_cast(msg_type) + << " on control rail " << control_rail_id; + control_rails_[control_rail_id]->releaseRequest(req); + return status; + } + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::progressActiveDataRails() { + std::vector rails_to_process; + + // Copy active rails under lock to avoid iterator invalidation + { + std::lock_guard lock(active_rails_mutex_); + if (active_rails_.empty()) { + return NIXL_IN_PROG; // No active rails to process + } + rails_to_process.assign(active_rails_.begin(), active_rails_.end()); + } + + // Process rails without holding the lock + bool any_completions = false; + + for (size_t rail_id : rails_to_process) { + if (rail_id >= data_rails_.size()) { + NIXL_ERROR << "Invalid active rail ID: " << rail_id; + continue; + } + // Process completions on active data rails + nixl_status_t status = data_rails_[rail_id]->progressCompletionQueue(false); + if (status == NIXL_SUCCESS) { + any_completions = true; + NIXL_DEBUG << "Processed completions on active data rail " << rail_id; + } else if (status != NIXL_IN_PROG && status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to process completions on active data rail " << rail_id; + // Continue processing other active rails even if one fails + } + } + + if (any_completions) { + NIXL_TRACE << "Processed " << rails_to_process.size() << " active rails, completions found"; + } + + return any_completions ? NIXL_SUCCESS : NIXL_IN_PROG; +} + +nixl_status_t +nixlLibfabricRailManager::progressAllControlRails() { + bool any_completions = false; + for (size_t rail_id = 0; rail_id < num_control_rails_; ++rail_id) { + nixl_status_t status = + control_rails_[rail_id]->progressCompletionQueue(true); // Blocking for control rails + if (status == NIXL_SUCCESS) { + any_completions = true; + NIXL_DEBUG << "Processed completion on control rail " << rail_id; + } else if (status != NIXL_IN_PROG && status != NIXL_SUCCESS) { + any_completions = true; + NIXL_ERROR << "Failed to process completion on control rail " << rail_id; + return NIXL_ERR_BACKEND; + } + } + return any_completions ? NIXL_SUCCESS : NIXL_IN_PROG; +} + +nixl_status_t +nixlLibfabricRailManager::validateAllRailsInitialized() { + for (size_t rail_id = 0; rail_id < data_rails_.size(); ++rail_id) { + if (!data_rails_[rail_id]->isProperlyInitialized()) { + NIXL_ERROR << "Rail " << rail_id << " is not properly initialized"; + return NIXL_ERR_BACKEND; + } + } + NIXL_DEBUG << "All " << data_rails_.size() << " rails are properly initialized"; + return NIXL_SUCCESS; +} + +struct fid_mr * +nixlLibfabricRailManager::getMemoryDescriptor(size_t rail_id, struct fid_mr *mr) { + if (rail_id >= data_rails_.size()) { + NIXL_ERROR << "Invalid rail index " << rail_id; + return nullptr; + } + return static_cast(data_rails_[rail_id]->getMemoryDescriptor(mr)); +} + +nixl_status_t +nixlLibfabricRailManager::serializeMemoryKeys(const std::vector &keys, + void *buffer, + std::string &str) const { + nixlSerDes ser_des; + // Serialize all rail keys instead of just the first one + for (size_t rail_id = 0; rail_id < keys.size(); ++rail_id) { + std::string key_name = "key_" + std::to_string(rail_id); + ser_des.addBuf(key_name.c_str(), &keys[rail_id], sizeof(keys[rail_id])); + } + + ser_des.addBuf("addr", &buffer, sizeof(buffer)); + str = ser_des.exportStr(); + + NIXL_DEBUG << "Serialized memory keys for " << keys.size() << " rails" << " (buffer: " << buffer + << ", size: " << str.length() << " bytes)"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::deserializeMemoryKeys(const std::string &serialized_data, + std::vector &keys_out, + uint64_t &remote_addr_out) const { + nixlSerDes ser_des; + ser_des.importStr(serialized_data); + // Load all rail keys instead of just one + keys_out.clear(); + keys_out.reserve(data_rails_.size()); + for (size_t rail_id = 0; rail_id < data_rails_.size(); ++rail_id) { + std::string key_name = "key_" + std::to_string(rail_id); + uint64_t remote_key; + nixl_status_t status = ser_des.getBuf(key_name.c_str(), &remote_key, sizeof(remote_key)); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to get key " << key_name << " for rail " << rail_id; + return NIXL_ERR_BACKEND; + } + keys_out.push_back(remote_key); + } + nixl_status_t addr_status = ser_des.getBuf("addr", &remote_addr_out, sizeof(remote_addr_out)); + if (addr_status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to get remote address"; + return NIXL_ERR_BACKEND; + } + NIXL_DEBUG << "Deserialized memory keys for " << keys_out.size() << " rails" + << " (remote addr: " << (void *)remote_addr_out << ")"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::serializeConnectionInfo(const std::string &user_prefix, + std::string &str) const { + + nixlSerDes ser_des; + + // Use user prefix with standard suffixes + std::string data_prefix = user_prefix + "_data_ep_"; + std::string control_prefix = user_prefix + "_control_ep_"; + + serializeRailEndpoints(ser_des, data_prefix, RailType::DATA); + serializeRailEndpoints(ser_des, control_prefix, RailType::CONTROL); + str = ser_des.exportStr(); + NIXL_DEBUG << "Connection info serialized with prefix " << user_prefix + << ", size: " << str.length(); + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricRailManager::deserializeConnectionInfo( + const std::string &user_prefix, + const std::string &serialized_data, + std::vector> &data_endpoints_out, + std::vector> &control_endpoints_out) const { + + nixlSerDes ser_des; + ser_des.importStr(serialized_data); + + // Use user prefix with standard suffixes + std::string data_prefix = user_prefix + "_data_ep_"; + std::string control_prefix = user_prefix + "_control_ep_"; + nixl_status_t data_status = + deserializeRailEndpoints(ser_des, data_prefix, data_rails_.size(), data_endpoints_out); + if (data_status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to deserialize data rail endpoints with prefix: " << data_prefix; + return data_status; + } + nixl_status_t control_status = deserializeRailEndpoints( + ser_des, control_prefix, control_rails_.size(), control_endpoints_out); + if (control_status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to deserialize control rail endpoints with prefix: " + << control_prefix; + return control_status; + } + + NIXL_DEBUG << "Connection info deserialized with prefix " << user_prefix << ": " + << data_endpoints_out.size() << " data endpoints, " << control_endpoints_out.size() + << " control endpoints"; + return NIXL_SUCCESS; +} + +void +nixlLibfabricRailManager::serializeRailEndpoints(nixlSerDes &ser_des, + const std::string &key_prefix, + RailType rail_type) const { + auto &rails = (rail_type == RailType::DATA) ? data_rails_ : control_rails_; + const char *rail_type_str = (rail_type == RailType::DATA) ? "data" : "control"; + + for (size_t rail_id = 0; rail_id < rails.size(); ++rail_id) { + std::string rail_key = key_prefix + std::to_string(rail_id); + const char *ep_name = rails[rail_id]->ep_name; + size_t ep_name_len = sizeof(rails[rail_id]->ep_name); + + ser_des.addBuf(rail_key.c_str(), ep_name, ep_name_len); + } + + NIXL_DEBUG << "Serialized " << rails.size() << " " << rail_type_str << " rail endpoints"; +} + +nixl_status_t +nixlLibfabricRailManager::deserializeRailEndpoints( + nixlSerDes &ser_des, + const std::string &key_prefix, + size_t expected_count, + std::vector> &endpoints_out) const { + endpoints_out.resize(expected_count); + + for (size_t rail_id = 0; rail_id < expected_count; ++rail_id) { + std::string rail_key = key_prefix + std::to_string(rail_id); + + // First check if the key exists and get its length + ssize_t actual_len = ser_des.getBufLen(rail_key); + if (actual_len <= 0) { + NIXL_ERROR << "Key " << rail_key << " not found or has invalid length: " << actual_len; + return NIXL_ERR_BACKEND; + } + + if (actual_len > (ssize_t)endpoints_out[rail_id].size()) { + NIXL_ERROR << "Buffer too small for rail " << rail_id << ", need " << actual_len + << " bytes, have " << endpoints_out[rail_id].size(); + return NIXL_ERR_BACKEND; + } + + // Get the actual data + nixl_status_t status = ser_des.getBuf(rail_key, endpoints_out[rail_id].data(), actual_len); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to get endpoint address for rail " << rail_id + << " with key: " << rail_key << ", status: " << status; + return NIXL_ERR_BACKEND; + } + } + + NIXL_DEBUG << "Successfully deserialized " << expected_count << " rail endpoints"; + return NIXL_SUCCESS; +} + +void +nixlLibfabricRailManager::markRailActive(size_t rail_id) { + if (rail_id >= data_rails_.size()) { + NIXL_ERROR << "Invalid rail ID for markRailActive: " << rail_id; + return; + } + + std::lock_guard lock(active_rails_mutex_); + bool was_inserted = active_rails_.insert(rail_id).second; + + if (was_inserted) { + NIXL_DEBUG << "Marked rail " << rail_id + << " as active (total active: " << active_rails_.size() << ")"; + } else { + NIXL_TRACE << "Rail " << rail_id << " was already active"; + } +} + +void +nixlLibfabricRailManager::markRailInactive(size_t rail_id) { + std::lock_guard lock(active_rails_mutex_); + size_t erased = active_rails_.erase(rail_id); + if (erased > 0) { + NIXL_DEBUG << "Marked rail " << rail_id + << " as inactive (total active: " << active_rails_.size() << ")"; + } else { + NIXL_TRACE << "Rail " << rail_id << " was not in active set"; + } +} + +void +nixlLibfabricRailManager::clearActiveRails() { + std::lock_guard lock(active_rails_mutex_); + size_t cleared_count = active_rails_.size(); + active_rails_.clear(); + NIXL_DEBUG << "Cleared " << cleared_count << " active rails"; +} + +size_t +nixlLibfabricRailManager::getActiveRailCount() const { + std::lock_guard lock(active_rails_mutex_); + return active_rails_.size(); +} diff --git a/src/utils/libfabric/libfabric_rail_manager.h b/src/utils/libfabric/libfabric_rail_manager.h new file mode 100644 index 0000000000..e41e29a1d1 --- /dev/null +++ b/src/utils/libfabric/libfabric_rail_manager.h @@ -0,0 +1,328 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_MANAGER_H +#define NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_MANAGER_H + +#include +#include +#include +#include +#include +#include +#include +#include "libfabric_rail.h" + +#ifdef HAVE_CUDA +#include +#include +#endif + +// Forward declarations +class nixlLibfabricTopology; + +/** Central manager for multi-rail RDMA operations with topology awareness */ +class nixlLibfabricRailManager { +public: + /** Initialize rail manager with topology discovery and create rails based on available EFA + * devices + * @param striping_threshold Size threshold for enabling multi-rail striping + * @throws std::runtime_error if topology discovery or rail creation fails + */ + nixlLibfabricRailManager(size_t striping_threshold); + /** Destroy rail manager and cleanup all resources */ + ~nixlLibfabricRailManager(); + + // Rail management + /** Create data rails for high-bandwidth transfers (one per EFA device) + * @param efa_devices List of EFA device names to create rails on + * @param provider_name Provider name ("efa" or "efa-direct") + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + createDataRails(const std::vector &efa_devices, const std::string &provider_name); + + /** Create control rails for connection management and notifications + * @param efa_devices List of EFA device names + * @param provider_name Provider name ("efa" or "efa-direct") + * @param num_control_rails Number of control rails to create + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + createControlRails(const std::vector &efa_devices, + const std::string &provider_name, + size_t num_control_rails); + + // Access rails + /** Get reference to data rail by ID */ + nixlLibfabricRail & + getDataRail(size_t rail_id) { + return *data_rails_[rail_id]; + } + + /** Get const reference to data rail by ID */ + const nixlLibfabricRail & + getDataRail(size_t rail_id) const { + return *data_rails_[rail_id]; + } + + /** Get reference to control rail by ID */ + nixlLibfabricRail & + getControlRail(size_t rail_id) { + return *control_rails_[rail_id]; + } + + /** Get const reference to control rail by ID */ + const nixlLibfabricRail & + getControlRail(size_t rail_id) const { + return *control_rails_[rail_id]; + } + + /** Get total number of data rails */ + size_t + getNumDataRails() const { + return data_rails_.size(); + } + + /** Get total number of control rails */ + size_t + getNumControlRails() const { + return control_rails_.size(); + } + + // Memory registration management + /** Register memory with topology-aware rail selection based on memory type and location + * @param buffer Memory buffer to register + * @param length Buffer size in bytes + * @param mem_type Memory type (DRAM_SEG or VRAM_SEG) + * @param gpu_id GPU device ID (used for VRAM_SEG, ignored for DRAM_SEG) + * @param mr_list_out Memory registration handles, indexed by rail ID + * @param key_list_out Remote access keys, indexed by rail ID + * @param selected_rails_out List of rail IDs where memory was registered + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + registerMemory(void *buffer, + size_t length, + nixl_mem_t mem_type, + int gpu_id, + std::vector &mr_list_out, + std::vector &key_list_out, + std::vector &selected_rails_out); + /** Deregister memory from specified rails + * @param selected_rails List of rail IDs to deregister from + * @param mr_list Memory registration handles to deregister + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + deregisterMemory(const std::vector &selected_rails, + const std::vector &mr_list); + + // Connection Management APIs + /** Rail type enumeration for connection operations */ + enum class RailType { DATA, CONTROL }; + /** Insert addresses into address vectors for all rails of specified type + * @param rail_type Type of rails to operate on (DATA or CONTROL) + * @param endpoints Remote endpoint addresses to insert + * @param fi_addrs_out Libfabric address handles for inserted endpoints + * @param ep_names_out Local endpoint names for reference + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + insertAllAddresses(RailType rail_type, + const std::vector> &endpoints, + std::vector &fi_addrs_out, + std::vector &ep_names_out); + /** Clean up connection resources for specified rail type + * @param rail_type Type of rails to clean up (DATA or CONTROL) + * @param fi_addrs_to_remove Libfabric addresses to remove + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + cleanupConnection(RailType rail_type, const std::vector &fi_addrs_to_remove); + + /** Single-pass transfer preparation and submission with automatic striping/round-robin + * @param op_type Operation type (WRITE or READ) + * @param local_addr Local memory address + * @param transfer_size Total transfer size + * @param remote_base_addr Remote memory base address + * @param selected_rails Rails to use for the transfer + * @param local_mrs Local memory registrations + * @param remote_keys Remote access keys + * @param dest_addrs Destination addresses for each rail + * @param agent_idx Remote agent index for immediate data + * @param completion_callback Callback for completion notification + * @param binary_notif Binary notification to populate with XFER_IDs + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + prepareAndSubmitTransfer(nixlLibfabricReq::OpType op_type, + void *local_addr, + size_t transfer_size, + uint64_t remote_base_addr, + const std::vector &selected_rails, + const std::vector &local_mrs, + const std::vector &remote_keys, + const std::vector &dest_addrs, + uint16_t agent_idx, + std::function completion_callback, + BinaryNotification *binary_notif); + /** Determine if striping should be used for given transfer size + * @param transfer_size Size of the transfer in bytes + * @return true if striping should be used, false for round-robin + */ + bool + shouldUseStriping(size_t transfer_size) const; + + // Control Message APIs + /** Control message types for rail communication */ + enum class ControlMessageType { + NOTIFICATION, ///< User notification message + CONNECTION_REQ, ///< Connection establishment request + CONNECTION_ACK, ///< Connection acknowledgment + DISCONNECT_REQ, ///< Disconnection request + }; + /** Send control message via control rail + * @param msg_type Type of control message + * @param req Control request with data buffer + * @param dest_addr Destination address + * @param agent_idx Agent index for message routing + * @param completion_callback Optional completion callback + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + postControlMessage(ControlMessageType msg_type, + nixlLibfabricReq *req, + fi_addr_t dest_addr, + uint16_t agent_idx = 0, + std::function completion_callback = nullptr); + // Progress APIs + /** Process completions on active data rails only (optimized for CPU overhead) + * @return NIXL_SUCCESS if completions processed, NIXL_IN_PROG if none, error on failure + */ + nixl_status_t + progressActiveDataRails(); + /** Process completions on all control rails for connection management and notifications + * @return NIXL_SUCCESS if completions processed, NIXL_IN_PROG if none, error on failure + */ + nixl_status_t + progressAllControlRails(); + /** Validate that all rails are properly initialized + * @return NIXL_SUCCESS if all rails initialized, error code otherwise + */ + nixl_status_t + validateAllRailsInitialized(); + + // Active Rail Management APIs + /** Mark rail as active for progress tracking optimization */ + void + markRailActive(size_t rail_id); + + /** Mark rail as inactive for progress tracking optimization */ + void + markRailInactive(size_t rail_id); + + /** Clear all active rail markings */ + void + clearActiveRails(); + + /** Get count of currently active rails */ + size_t + getActiveRailCount() const; + + // Memory Descriptor APIs + /** Get memory descriptor for specified rail and MR */ + struct fid_mr * + getMemoryDescriptor(size_t rail_id, struct fid_mr *mr); + + // SerDes-based Memory Key Serialization + /** Serialize memory keys and buffer address for remote access + * @param keys Remote access keys for all rails + * @param buffer Memory buffer address + * @param str Serialized data string + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + serializeMemoryKeys(const std::vector &keys, void *buffer, std::string &str) const; + /** Deserialize memory keys and remote address + * @param serialized_data Serialized memory information + * @param keys_out Remote access keys for all rails + * @param remote_addr_out Remote buffer address + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + deserializeMemoryKeys(const std::string &serialized_data, + std::vector &keys_out, + uint64_t &remote_addr_out) const; + // SerDes-based Connection Info Serialization + /** Serialize connection information for all rails + * @param user_prefix Prefix for serialization keys (e.g., "src" or "dest") + * @param str Serialized connection information + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + serializeConnectionInfo(const std::string &user_prefix, std::string &str) const; + /** Deserialize connection information for all rails + * @param user_prefix Prefix used during serialization + * @param serialized_data Serialized connection information + * @param data_endpoints_out Data rail endpoint addresses + * @param control_endpoints_out Control rail endpoint addresses + * @return NIXL_SUCCESS on success, error code on failure + */ + nixl_status_t + deserializeConnectionInfo( + const std::string &user_prefix, + const std::string &serialized_data, + std::vector> &data_endpoints_out, + std::vector> &control_endpoints_out) const; + +private: + size_t striping_threshold_; + // Rail allocation + std::vector> data_rails_; + std::vector> control_rails_; + + size_t num_data_rails_; + size_t num_control_rails_; + + std::unique_ptr topology; + + // EFA device to rail mapping + std::unordered_map efa_device_to_rail_map; + + // Active Rail Tracking System + std::unordered_set active_rails_; + mutable std::mutex active_rails_mutex_; + + // Internal rail selection method + std::vector + selectRailsForMemory(void *mem_addr, nixl_mem_t mem_type, int gpu_id) const; + + // Helper functions for connection SerDes + void + serializeRailEndpoints(nixlSerDes &ser_des, + const std::string &key_prefix, + RailType rail_type) const; + nixl_status_t + deserializeRailEndpoints( + nixlSerDes &ser_des, + const std::string &key_prefix, + size_t expected_count, + std::vector> &endpoints_out) const; +}; + +#endif // NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_RAIL_MANAGER_H diff --git a/src/utils/libfabric/libfabric_topology.cpp b/src/utils/libfabric/libfabric_topology.cpp new file mode 100644 index 0000000000..e8667408ce --- /dev/null +++ b/src/utils/libfabric/libfabric_topology.cpp @@ -0,0 +1,667 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric_topology.h" +#include "libfabric_common.h" +#include "common/nixl_log.h" + +#include +#include +#include +#include + +#include +#include + +#ifdef HAVE_CUDA +#include +#endif + +nixlLibfabricTopology::nixlLibfabricTopology() + : num_gpus(0), + num_numa_nodes(0), + num_efa_devices(0), + topology_discovered(false), + hwloc_topology(nullptr) { + + NIXL_TRACE << "Starting automatic topology discovery"; + + // Discover topology immediately - hard error if it fails + nixl_status_t status = discoverTopology(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Topology discovery failed - this is a fatal error"; + throw std::runtime_error( + "Failed to discover system topology - cannot proceed without topology information"); + } + NIXL_TRACE << "Topology discovery completed successfully"; + printTopologyInfo(); +} + +nixlLibfabricTopology::~nixlLibfabricTopology() { + cleanupHwlocTopology(); +} + +nixl_status_t +nixlLibfabricTopology::discoverTopology() { + NIXL_TRACE << "Starting hwloc-based topology discovery"; + // Initialize hwloc topology + nixl_status_t status = initHwlocTopology(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to initialize hwloc topology"; + return status; + } + // Discover EFA devices using libfabric + status = discoverEfaDevices(); + if (status != NIXL_SUCCESS) { + return status; + } + // Build PCIe to Libfabric device mapping + status = buildPcieToLibfabricMapping(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to build PCIe to Libfabric mapping - this is required for topology " + "discovery"; + return status; + } + // Discover hardware topology using hwloc + status = discoverHwlocTopology(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to discover hwloc topology"; + return status; + } + // Build GPU to EFA mapping based on PCIe topology + status = buildGpuToEfaMapping(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to build GPU to EFA mapping"; + return status; + } + topology_discovered = true; + NIXL_TRACE << "Topology discovery completed successfully"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::discoverEfaDevices() { + // Use the utility function from libfabric_common + auto efa_result = LibfabricUtils::getAvailableEfaDevices(); + efa_fabric_name = efa_result.first; + all_efa_devices = efa_result.second; + + num_efa_devices = all_efa_devices.size(); + if (all_efa_devices.empty()) { + NIXL_ERROR << "No EFA devices found"; + return NIXL_ERR_BACKEND; + } + NIXL_TRACE << "Discovered " << num_efa_devices << " EFA devices"; + for (size_t i = 0; i < all_efa_devices.size(); ++i) { + NIXL_TRACE << "EFA device " << i << ": " << all_efa_devices[i]; + } + return NIXL_SUCCESS; +} + +std::vector +nixlLibfabricTopology::getEfaDevicesForGpu(int gpu_id) const { + auto it = gpu_to_efa_devices.find(gpu_id); + if (it != gpu_to_efa_devices.end()) { + return it->second; + } + NIXL_WARN << "No EFA devices found for GPU " << gpu_id << ", returning all devices"; + return all_efa_devices; +} + +bool +nixlLibfabricTopology::isValidGpuId(int gpu_id) const { + return gpu_id >= 0 && gpu_id < num_gpus; +} + +bool +nixlLibfabricTopology::isValidEfaDevice(const std::string &efa_device) const { + return std::find(all_efa_devices.begin(), all_efa_devices.end(), efa_device) != + all_efa_devices.end(); +} + +void +nixlLibfabricTopology::printTopologyInfo() const { + NIXL_TRACE << "=== Libfabric Topology Information ==="; + NIXL_TRACE << "Topology discovered: " << (topology_discovered ? "Yes" : "No"); + NIXL_TRACE << "Number of GPUs: " << num_gpus; + NIXL_TRACE << "Number of NUMA nodes: " << num_numa_nodes; + NIXL_TRACE << "Number of EFA devices: " << num_efa_devices; + NIXL_TRACE << "EFA devices: "; + for (size_t i = 0; i < all_efa_devices.size(); ++i) { + NIXL_INFO << " [" << i << "] " << all_efa_devices[i]; + } + NIXL_TRACE << "GPU → EFA mapping:"; + for (const auto &pair : gpu_to_efa_devices) { + std::stringstream ss; + ss << " GPU " << pair.first << " → ["; + for (size_t i = 0; i < pair.second.size(); ++i) { + if (i > 0) ss << ", "; + ss << pair.second[i]; + } + ss << "]"; + NIXL_INFO << ss.str(); + } + NIXL_TRACE << "Host memory (DRAM) will use all available EFA devices for maximum bandwidth"; + NIXL_TRACE << "====================================="; +} + +std::string +nixlLibfabricTopology::getTopologyString() const { + std::stringstream ss; + ss << "Libfabric Topology: "; + ss << "GPUs=" << num_gpus << ", "; + ss << "NUMA=" << num_numa_nodes << ", "; + ss << "EFA=" << num_efa_devices << ", "; + ss << "Discovered=" << (topology_discovered ? "Yes" : "No"); + return ss.str(); +} + +// hwloc-based implementation methods + +nixl_status_t +nixlLibfabricTopology::initHwlocTopology() { + if (hwloc_topology) { + cleanupHwlocTopology(); + } + int ret = hwloc_topology_init(&hwloc_topology); + if (ret != 0) { + NIXL_ERROR << "Failed to initialize hwloc topology: " << ret; + return NIXL_ERR_BACKEND; + } + // Enable I/O device discovery - this is the key to seeing EFA devices! +#if (HWLOC_API_VERSION >= 0x00020000) + enum hwloc_type_filter_e filter = HWLOC_TYPE_FILTER_KEEP_ALL; + ret = hwloc_topology_set_io_types_filter(hwloc_topology, filter); + if (ret != 0) { + NIXL_WARN << "Failed to set IO types filter: " << ret << ", continuing anyway"; + } +#else + unsigned long flags = hwloc_topology_get_flags(hwloc_topology); + flags |= HWLOC_TOPOLOGY_FLAG_WHOLE_IO; + ret = hwloc_topology_set_flags(hwloc_topology, flags); + if (ret != 0) { + NIXL_WARN << "Failed to set WHOLE_IO flag: " << ret << ", continuing anyway"; + } +#endif + ret = hwloc_topology_load(hwloc_topology); + if (ret != 0) { + NIXL_ERROR << "Failed to load hwloc topology: " << ret; + hwloc_topology_destroy(hwloc_topology); + hwloc_topology = nullptr; + return NIXL_ERR_BACKEND; + } + NIXL_TRACE << "hwloc topology initialized successfully with IO device support"; + return NIXL_SUCCESS; +} + +void +nixlLibfabricTopology::cleanupHwlocTopology() { + if (hwloc_topology) { + hwloc_topology_destroy(hwloc_topology); + hwloc_topology = nullptr; + } +} + +nixl_status_t +nixlLibfabricTopology::discoverHwlocTopology() { + if (!hwloc_topology) { + NIXL_ERROR << "hwloc topology not initialized"; + return NIXL_ERR_BACKEND; + } + // Discover GPUs and EFA devices using hwloc + nixl_status_t status = discoverGpusWithHwloc(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to discover GPUs with hwloc"; + return status; + } + status = discoverEfaDevicesWithHwloc(); + if (status != NIXL_SUCCESS) { + NIXL_ERROR << "Failed to discover EFA devices with hwloc"; + return status; + } + // Discover NUMA topology + num_numa_nodes = hwloc_get_nbobjs_by_type(hwloc_topology, HWLOC_OBJ_NUMANODE); + if (num_numa_nodes == 0) { + num_numa_nodes = 1; // Fallback to single NUMA node + } + NIXL_TRACE << "Discovered " << num_gpus << " GPUs and " << num_numa_nodes + << " NUMA nodes via hwloc"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::discoverGpusWithHwloc() { + num_gpus = 0; + // Find all PCI devices and log detailed information + hwloc_obj_t pci_obj = nullptr; + while ((pci_obj = hwloc_get_next_pcidev(hwloc_topology, pci_obj)) != nullptr) { + if (isNvidiaGpu(pci_obj)) { + std::string pcie_addr = getPcieAddressFromHwlocObj(pci_obj); + // Get device and vendor info + uint16_t vendor_id = pci_obj->attr->pcidev.vendor_id; + uint16_t device_id = pci_obj->attr->pcidev.device_id; + uint16_t class_id = pci_obj->attr->pcidev.class_id; + + NIXL_TRACE << "Found NVIDIA GPU " << num_gpus << ": " << pcie_addr << " (vendor=0x" + << std::hex << vendor_id << ", device=0x" << device_id << ", class=0x" + << class_id << std::dec << ")"; + + num_gpus++; + } + } + + NIXL_TRACE << "Discovered " << num_gpus << " NVIDIA GPUs via hwloc"; + + // If we found more than 8 GPUs on P5en, investigate further + if (num_gpus > 8) { + NIXL_WARN << "Found " << num_gpus + << " NVIDIA GPUs, but P5en should have 8. Investigating..."; + + // List all NVIDIA devices to understand what we're seeing + pci_obj = nullptr; + int gpu_count = 0; + while ((pci_obj = hwloc_get_next_pcidev(hwloc_topology, pci_obj)) != nullptr) { + if (pci_obj->attr->pcidev.vendor_id == 0x10de) { // NVIDIA + std::string pcie_addr = getPcieAddressFromHwlocObj(pci_obj); + uint16_t device_id = pci_obj->attr->pcidev.device_id; + uint16_t class_id = pci_obj->attr->pcidev.class_id; + + NIXL_WARN << "NVIDIA device " << gpu_count << ": " << pcie_addr << " (device=0x" + << std::hex << device_id << ", class=0x" << class_id << std::dec << ")"; + gpu_count++; + } + } + } + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::discoverEfaDevicesWithHwloc() { + // EFA devices are already discovered via libfabric + // This method validates the hwloc discovery matches libfabric discovery + int hwloc_efa_count = 0; + hwloc_obj_t pci_obj = nullptr; + while ((pci_obj = hwloc_get_next_pcidev(hwloc_topology, pci_obj)) != nullptr) { + if (isEfaDevice(pci_obj)) { + hwloc_efa_count++; + NIXL_TRACE << "Found EFA device via hwloc: " << getPcieAddressFromHwlocObj(pci_obj); + } + } + + NIXL_TRACE << "hwloc found " << hwloc_efa_count << " EFA devices, libfabric found " + << num_efa_devices; + + if (hwloc_efa_count != num_efa_devices) { + NIXL_WARN << "Mismatch between hwloc (" << hwloc_efa_count << ") and libfabric (" + << num_efa_devices << ") EFA device counts"; + } + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::buildPcieToLibfabricMapping() { + pcie_to_libfabric_map.clear(); + libfabric_to_pcie_map.clear(); + + // Get EFA device info with PCIe addresses from libfabric + struct fi_info *hints, *info; + + hints = fi_allocinfo(); + if (!hints) { + NIXL_ERROR << "Failed to allocate fi_info for PCIe mapping"; + return NIXL_ERR_BACKEND; + } + + hints->fabric_attr->prov_name = strdup("efa"); + int ret = fi_getinfo(FI_VERSION(1, 9), NULL, NULL, 0, hints, &info); + if (ret) { + NIXL_ERROR << "fi_getinfo failed for PCIe mapping: " << fi_strerror(-ret); + fi_freeinfo(hints); + return NIXL_ERR_BACKEND; + } + + for (struct fi_info *cur = info; cur; cur = cur->next) { + if (cur->domain_attr && cur->domain_attr->name && cur->nic && cur->nic->bus_attr) { + std::string libfabric_name = cur->domain_attr->name; + // Extract PCIe address from bus_attr if available + if (cur->nic->bus_attr->bus_type == FI_BUS_PCI && + cur->nic->bus_attr->attr.pci.domain_id != FI_ADDR_UNSPEC) { + char pcie_addr[32]; + snprintf(pcie_addr, + sizeof(pcie_addr), + "%x:%02x:%02x.%x", + cur->nic->bus_attr->attr.pci.domain_id, + cur->nic->bus_attr->attr.pci.bus_id, + cur->nic->bus_attr->attr.pci.device_id, + cur->nic->bus_attr->attr.pci.function_id); + + std::string pcie_address = pcie_addr; + pcie_to_libfabric_map[pcie_address] = libfabric_name; + libfabric_to_pcie_map[libfabric_name] = pcie_address; + + NIXL_TRACE << "Mapped PCIe " << pcie_address << " → Libfabric " << libfabric_name; + } + } + } + + fi_freeinfo(info); + fi_freeinfo(hints); + NIXL_TRACE << "Built PCIe to Libfabric mapping for " << pcie_to_libfabric_map.size() + << " devices"; + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::buildGpuToEfaMapping() { + gpu_to_efa_devices.clear(); + // Implement NIXL's topology-aware GPU-EFA grouping algorithm + nixl_status_t status = buildTopologyAwareGrouping(); + if (status != NIXL_SUCCESS) { + NIXL_WARN << "Topology-aware grouping failed, using fallback"; + return buildFallbackMapping(); + } + + NIXL_TRACE << "Built GPU→EFA mapping for " << gpu_to_efa_devices.size() + << " GPUs using topology-aware algorithm"; + + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::buildTopologyAwareGrouping() { + // Step 1: Build NIC info structures by correlating libfabric with hwloc + std::vector discovered_nics; + std::vector discovered_gpus; + // Discover NICs by correlating libfabric devices with hwloc objects + for (const auto &pair : pcie_to_libfabric_map) { + const std::string &pcie_addr = pair.first; + const std::string &libfabric_name = pair.second; + + // Parse PCIe address + uint16_t domain_id; + uint8_t bus_id, device_id, function_id; + if (sscanf(pcie_addr.c_str(), + "%hx:%hhx:%hhx.%hhx", + &domain_id, + &bus_id, + &device_id, + &function_id) != 4) { + NIXL_WARN << "Failed to parse PCIe address: " << pcie_addr; + continue; + } + + // Find corresponding hwloc object + hwloc_obj_t hwloc_node = + hwloc_get_pcidev_by_busid(hwloc_topology, domain_id, bus_id, device_id, function_id); + + if (hwloc_node) { + NicInfo nic; + nic.libfabric_name = libfabric_name; + nic.hwloc_node = hwloc_node; + nic.domain_id = domain_id; + nic.bus_id = bus_id; + nic.device_id = device_id; + nic.function_id = function_id; + discovered_nics.push_back(nic); + NIXL_TRACE << "Correlated NIC: " << pcie_addr << " → " << libfabric_name; + } else { + NIXL_WARN << "Could not find hwloc object for PCIe address: " << pcie_addr; + } + } + // Step 2: Discover GPUs + hwloc_obj_t pci_obj = nullptr; + while ((pci_obj = hwloc_get_next_pcidev(hwloc_topology, pci_obj)) != nullptr) { + if (isNvidiaGpu(pci_obj)) { + GpuInfo gpu; + gpu.hwloc_node = pci_obj; + gpu.domain_id = pci_obj->attr->pcidev.domain; + gpu.bus_id = pci_obj->attr->pcidev.bus; + gpu.device_id = pci_obj->attr->pcidev.dev; + gpu.function_id = pci_obj->attr->pcidev.func; + discovered_gpus.push_back(gpu); + } + } + + NIXL_TRACE << "Discovered " << discovered_nics.size() << " NICs and " << discovered_gpus.size() + << " GPUs for grouping"; + + if (discovered_nics.empty() || discovered_gpus.empty()) { + NIXL_WARN << "No NICs or GPUs found for grouping"; + return NIXL_ERR_BACKEND; + } + // Step 3: Implement NIXL's topology-aware grouping algorithm + std::vector nic_groups; + nixl_status_t status = groupNicsWithGpus(discovered_nics, discovered_gpus, nic_groups); + if (status != NIXL_SUCCESS) { + return status; + } + // Step 4: Convert groups to GPU→EFA mapping + for (size_t group_idx = 0; group_idx < nic_groups.size(); ++group_idx) { + const auto &group = nic_groups[group_idx]; + if (group.has_gpu) { + std::vector gpu_efa_devices; + for (const auto &nic : group.nics) { + gpu_efa_devices.push_back(nic.libfabric_name); + } + // Find GPU index in our discovered GPUs list + int gpu_index = -1; + for (size_t i = 0; i < discovered_gpus.size(); ++i) { + const auto &gpu = discovered_gpus[i]; + if (gpu.domain_id == group.closest_gpu.domain_id && + gpu.bus_id == group.closest_gpu.bus_id && + gpu.device_id == group.closest_gpu.device_id && + gpu.function_id == group.closest_gpu.function_id) { + gpu_index = static_cast(i); + break; + } + } + + if (gpu_index >= 0) { + gpu_to_efa_devices[gpu_index] = gpu_efa_devices; + + NIXL_TRACE << "GPU " << gpu_index << " (" << std::hex << group.closest_gpu.domain_id + << ":" << static_cast(group.closest_gpu.bus_id) << ":" + << static_cast(group.closest_gpu.device_id) << "." + << static_cast(group.closest_gpu.function_id) << std::dec << ") → " + << gpu_efa_devices.size() << " EFA devices"; + } + } + } + return NIXL_SUCCESS; +} + +nixl_status_t +nixlLibfabricTopology::buildFallbackMapping() { + // Fallback: if specific mapping failed, use simple approach + gpu_to_efa_devices.clear(); + // Give all devices to all GPUs (not optimal but functional) + for (int gpu_id = 0; gpu_id < num_gpus; ++gpu_id) { + gpu_to_efa_devices[gpu_id] = all_efa_devices; + } + return NIXL_SUCCESS; +} + + +// hwloc helper methods + +std::string +nixlLibfabricTopology::getPcieAddressFromHwlocObj(hwloc_obj_t obj) const { + if (!obj || obj->type != HWLOC_OBJ_PCI_DEVICE) { + return ""; + } + char pcie_addr[32]; + snprintf(pcie_addr, + sizeof(pcie_addr), + "%x:%02x:%02x.%x", + obj->attr->pcidev.domain, + obj->attr->pcidev.bus, + obj->attr->pcidev.dev, + obj->attr->pcidev.func); + return std::string(pcie_addr); +} + +bool +nixlLibfabricTopology::isNvidiaGpu(hwloc_obj_t obj) const { + if (!obj || obj->type != HWLOC_OBJ_PCI_DEVICE) { + return false; + } + // NVIDIA vendor ID is 0x10de + if (obj->attr->pcidev.vendor_id != 0x10de) { + return false; + } + // Only count devices with GPU class (0x300-0x3ff for display controllers) + // Class 0x302 is 3D controller (GPU), 0x680 is other devices (network, etc.) + uint16_t class_id = obj->attr->pcidev.class_id; + return (class_id >= 0x300 && class_id < 0x400); +} + +bool +nixlLibfabricTopology::isEfaDevice(hwloc_obj_t obj) const { + if (!obj || obj->type != HWLOC_OBJ_PCI_DEVICE) { + return false; + } + + // Amazon EFA vendor ID is 0x1d0f, device ID can be 0xefa0, 0xefa1, or 0xefa2 + return obj->attr->pcidev.vendor_id == 0x1d0f && + (obj->attr->pcidev.device_id == 0xefa0 || obj->attr->pcidev.device_id == 0xefa1 || + obj->attr->pcidev.device_id == 0xefa2); +} + +nixl_status_t +nixlLibfabricTopology::groupNicsWithGpus(const std::vector &discovered_nics, + const std::vector &discovered_gpus, + std::vector &nic_groups) { + nic_groups.clear(); + + // Implement NIXL's topology-aware NIC grouping algorithm + + // Step 1: Mark topology nodes that have NICs in their subtree + std::map node_group_counts; + std::map> node_nics; + std::set nic_subtree_nodes; + // Mark all nodes that have NICs in their subtree and collect NICs per node + for (const auto &nic : discovered_nics) { + hwloc_obj_t node = nic.hwloc_node; + node_nics[node].push_back(nic); + while (node) { + nic_subtree_nodes.insert(node); + node = node->parent; + } + } + + // Step 2: For each GPU, walk up until finding a NIC subtree node and increment its count + std::map> node_gpus; + + for (const auto &gpu : discovered_gpus) { + hwloc_obj_t node = gpu.hwloc_node; + + while (node) { + if (nic_subtree_nodes.find(node) != nic_subtree_nodes.end()) { + node_group_counts[node]++; + node_gpus[node].push_back(gpu); + break; + } + node = node->parent; + } + } + + // Step 3: Collect all NICs that need to be grouped and assign them to ancestor nodes + std::map> ancestor_nics; + + for (const auto &pair : node_nics) { + hwloc_obj_t nic_node = pair.first; + const std::vector &nics = pair.second; + + // Find the ancestor with group count > 0 + hwloc_obj_t target_node = nic_node; + while (target_node) { + if (node_group_counts[target_node] > 0) { + // Add these NICs to this ancestor + ancestor_nics[target_node].insert( + ancestor_nics[target_node].end(), nics.begin(), nics.end()); + break; + } + target_node = target_node->parent; + } + // If no ancestor found with groups, create individual groups + if (!target_node) { + for (const auto &nic : nics) { + NicGroup group; + group.nics.push_back(nic); + group.has_gpu = false; + group.closest_gpu.hwloc_node = nullptr; + group.common_ancestor = nic.hwloc_node; + nic_groups.push_back(group); + } + } + } + // Step 4: Split NICs among GPUs for each ancestor node + for (const auto &pair : ancestor_nics) { + hwloc_obj_t ancestor = pair.first; + std::vector nics = pair.second; + int num_groups = node_group_counts[ancestor]; + const std::vector &gpus = node_gpus[ancestor]; + + if (num_groups > 0 && !gpus.empty()) { + // Sort NICs by bus ID for consistent assignment + std::sort(nics.begin(), nics.end(), [](const NicInfo &a, const NicInfo &b) { + if (a.bus_id != b.bus_id) return a.bus_id < b.bus_id; + return a.device_id < b.device_id; + }); + + // Split NICs among GPUs + int nics_per_group = nics.size() / num_groups; + int extra_nics = nics.size() % num_groups; + + size_t nic_idx = 0; + for (int group_idx = 0; group_idx < num_groups && group_idx < (int)gpus.size(); + ++group_idx) { + NicGroup group; + group.has_gpu = true; + group.closest_gpu = gpus[group_idx]; + group.common_ancestor = ancestor; + // Assign NICs to this group + int group_size = nics_per_group + (group_idx < extra_nics ? 1 : 0); + for (int i = 0; i < group_size && nic_idx < nics.size(); ++i, ++nic_idx) { + group.nics.push_back(nics[nic_idx]); + } + if (!group.nics.empty()) { + nic_groups.push_back(group); + } + } + } + } + + NIXL_TRACE << "NIXL topology grouping created " << nic_groups.size() << " NIC groups"; + + // Log the groups for debugging + for (size_t i = 0; i < nic_groups.size(); ++i) { + const auto &group = nic_groups[i]; + if (group.has_gpu) { + NIXL_TRACE << "Group " << i << ": GPU " << std::hex << group.closest_gpu.domain_id + << ":" << static_cast(group.closest_gpu.bus_id) << ":" + << static_cast(group.closest_gpu.device_id) << "." + << static_cast(group.closest_gpu.function_id) << std::dec << " → " + << group.nics.size() << " NICs"; + } else { + NIXL_TRACE << "Group " << i << ": No GPU → " << group.nics.size() << " NICs"; + } + } + return NIXL_SUCCESS; +} diff --git a/src/utils/libfabric/libfabric_topology.h b/src/utils/libfabric/libfabric_topology.h new file mode 100644 index 0000000000..e41509ef50 --- /dev/null +++ b/src/utils/libfabric/libfabric_topology.h @@ -0,0 +1,163 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_TOPOLOGY_H +#define NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_TOPOLOGY_H + +#include "libfabric_common.h" +#include "nixl.h" +#include +#include + +/** + * @brief Topology discovery and management for AWS instances with EFA devices + * + * Automatically discovers system topology using hwloc and maps GPUs to EFA devices + * based on PCIe proximity for optimal performance. Hard errors if topology discovery fails. + */ +class nixlLibfabricTopology { +private: + // GPU to EFA device mapping: GPU 0→[efa0,efa1], GPU 1→[efa2,efa3], etc. + std::map> gpu_to_efa_devices; + + // All available EFA devices discovered on this system + std::vector all_efa_devices; + + // EFA fabric name + std::string efa_fabric_name; + + // System information + int num_gpus; + int num_numa_nodes; + int num_efa_devices; + + // Discovery state + bool topology_discovered; + + // hwloc topology handle + hwloc_topology_t hwloc_topology; + + // PCIe to Libfabric device mapping + std::map pcie_to_libfabric_map; + std::map libfabric_to_pcie_map; + + // Helper methods + nixl_status_t + discoverEfaDevices(); + nixl_status_t + discoverTopology(); + + // hwloc-based discovery methods + nixl_status_t + initHwlocTopology(); + nixl_status_t + discoverHwlocTopology(); + nixl_status_t + buildPcieToLibfabricMapping(); + nixl_status_t + discoverGpusWithHwloc(); + nixl_status_t + discoverEfaDevicesWithHwloc(); + nixl_status_t + buildGpuToEfaMapping(); + void + cleanupHwlocTopology(); + + // Data structures for NIXL topology-aware grouping algorithm + struct NicInfo { + std::string libfabric_name; + hwloc_obj_t hwloc_node; + uint16_t domain_id; + uint8_t bus_id; + uint8_t device_id; + uint8_t function_id; + }; + + struct GpuInfo { + hwloc_obj_t hwloc_node; + uint16_t domain_id; + uint8_t bus_id; + uint8_t device_id; + uint8_t function_id; + }; + + struct NicGroup { + std::vector nics; + GpuInfo closest_gpu; + hwloc_obj_t common_ancestor; + bool has_gpu; + }; + + // NIXL topology-aware grouping algorithm methods + nixl_status_t + buildTopologyAwareGrouping(); + nixl_status_t + buildFallbackMapping(); + nixl_status_t + groupNicsWithGpus(const std::vector &discovered_nics, + const std::vector &discovered_gpus, + std::vector &nic_groups); + + // hwloc helper methods + std::string + getPcieAddressFromHwlocObj(hwloc_obj_t obj) const; + bool + isNvidiaGpu(hwloc_obj_t obj) const; + bool + isEfaDevice(hwloc_obj_t obj) const; + +public: + nixlLibfabricTopology(); // Automatically discovers topology, throws on failure + ~nixlLibfabricTopology(); + // GPU-based queries (main interface) + std::vector + getEfaDevicesForGpu(int gpu_id) const; + + // System information + int + getNumGpus() const { + return num_gpus; + } + + const std::vector & + getAllEfaDevices() const { + return all_efa_devices; + } + + const std::string & + getEFAfabricName() const { + return efa_fabric_name; + } + + // Validation + bool + isTopologyDiscovered() const { + return topology_discovered; + } + + bool + isValidGpuId(int gpu_id) const; + bool + isValidEfaDevice(const std::string &efa_device) const; + // Debug/info + void + printTopologyInfo() const; + std::string + getTopologyString() const; +}; + +#endif // NIXL_SRC_UTILS_LIBFABRIC_LIBFABRIC_TOPOLOGY_H diff --git a/src/utils/libfabric/meson.build b/src/utils/libfabric/meson.build new file mode 100644 index 0000000000..39fa98bca3 --- /dev/null +++ b/src/utils/libfabric/meson.build @@ -0,0 +1,69 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Source files +libfabric_utils_sources = files( + 'libfabric_rail.cpp', + 'libfabric_rail_manager.cpp', + 'libfabric_common.cpp', + 'libfabric_topology.cpp', + # More implementation files will be added as we create them +) + +# Header files +libfabric_utils_headers = files( + 'libfabric_rail.h', + 'libfabric_rail_manager.h', + 'libfabric_common.h', + 'libfabric_topology.h', +) + +# Find hwloc dependency +hwloc_dep = dependency('hwloc', required: true) + +# Set up dependencies and compile args +libfabric_utils_deps = [ + libfabric_dep, + hwloc_dep, + abseil_proj.get_variable('absl_log_dep'), +] + +libfabric_utils_cpp_args = [] + +# Add CUDA support if available +if cuda_dep.found() + libfabric_utils_deps += [cuda_dep] + libfabric_utils_cpp_args += ['-DHAVE_CUDA'] +endif + +# Create static library +libfabric_utils_lib = static_library( + 'nixl_libfabric_utils', + libfabric_utils_sources, + dependencies: libfabric_utils_deps, + cpp_args: libfabric_utils_cpp_args, + include_directories: [ + nixl_inc_dirs, + utils_inc_dirs, + ], + install: false, +) + +# Declare dependency for other parts of the build +libfabric_utils_dep = declare_dependency( + link_with: libfabric_utils_lib, + include_directories: include_directories('.'), +) diff --git a/src/utils/meson.build b/src/utils/meson.build index a60356f00e..12ee072479 100644 --- a/src/utils/meson.build +++ b/src/utils/meson.build @@ -18,3 +18,7 @@ subdir('serdes') subdir('ucx') subdir('stream') subdir('file') + +if libfabric_dep.found() + subdir('libfabric') +endif diff --git a/src/utils/serdes/meson.build b/src/utils/serdes/meson.build index 14c22d3895..1de69310a5 100644 --- a/src/utils/serdes/meson.build +++ b/src/utils/serdes/meson.build @@ -15,7 +15,8 @@ serdes_lib = library('serdes', 'serdes.cpp', 'serdes.h', - include_directories: nixl_inc_dirs, + dependencies: [absl_log_dep], + include_directories: [nixl_inc_dirs, utils_inc_dirs], install: true) serdes_interface = declare_dependency(link_with: serdes_lib) diff --git a/src/utils/serdes/serdes.cpp b/src/utils/serdes/serdes.cpp index 2c24fec6b9..b2f0bcfb25 100644 --- a/src/utils/serdes/serdes.cpp +++ b/src/utils/serdes/serdes.cpp @@ -15,6 +15,7 @@ * limitations under the License. */ #include "serdes.h" +#include "common/nixl_log.h" nixlSerDes::nixlSerDes() { workingStr = "nixlSerDes|"; @@ -48,8 +49,8 @@ nixl_status_t nixlSerDes::addStr(const std::string &tag, const std::string &str) std::string nixlSerDes::getStr(const std::string &tag){ if(workingStr.compare(des_offset, tag.size(), tag) != 0){ - //incorrect tag - return ""; + NIXL_ERROR << "Deserialization of tag " << tag << " failed"; + return ""; } ssize_t len; @@ -67,6 +68,8 @@ std::string nixlSerDes::getStr(const std::string &tag){ //move past string plus | delimiter des_offset += len + 1; + if (ret.empty()) NIXL_ERROR << "Deserialization of tag " << tag << " failed"; + return ret; } @@ -83,8 +86,8 @@ nixl_status_t nixlSerDes::addBuf(const std::string &tag, const void* buf, ssize_ ssize_t nixlSerDes::getBufLen(const std::string &tag) const{ if(workingStr.compare(des_offset, tag.size(), tag) != 0){ - //incorrect tag - return -1; + NIXL_ERROR << "Deserialization of tag " << tag << " failed"; + return -1; } ssize_t len; @@ -93,13 +96,15 @@ ssize_t nixlSerDes::getBufLen(const std::string &tag) const{ //_stringToBytes(&len, workingStr.data() + des_offset + tag.size(), sizeof(ssize_t)); _stringToBytes(&len, workingStr.substr(des_offset + tag.size(), sizeof(ssize_t)), sizeof(ssize_t)); + if (len == 0) NIXL_WARN << "In deserialization of tag " << tag << " the buffer length ios 0"; + return len; } nixl_status_t nixlSerDes::getBuf(const std::string &tag, void *buf, ssize_t len){ if(workingStr.compare(des_offset, tag.size(), tag) != 0){ - //incorrect tag - return NIXL_ERR_MISMATCH; + NIXL_ERROR << "Deserialization of tag " << tag << " failed"; + return NIXL_ERR_MISMATCH; } //skip over tag and size, which we assume has been read previously @@ -122,8 +127,8 @@ std::string nixlSerDes::exportStr() const { nixl_status_t nixlSerDes::importStr(const std::string &sdbuf) { if(sdbuf.compare(0, 11, "nixlSerDes|") != 0){ - //incorrect tag - return NIXL_ERR_MISMATCH; + NIXL_ERROR << "Deserialization failed, missing nixlSerDes tag"; + return NIXL_ERR_MISMATCH; } workingStr = sdbuf; diff --git a/src/utils/ucx/gpu_xfer_req_h.cpp b/src/utils/ucx/gpu_xfer_req_h.cpp new file mode 100644 index 0000000000..f33b8b9e5f --- /dev/null +++ b/src/utils/ucx/gpu_xfer_req_h.cpp @@ -0,0 +1,105 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "gpu_xfer_req_h.h" +#include "common/nixl_log.h" +#include "ucx_utils.h" +#include "rkey.h" +#include "config.h" + +extern "C" { +#ifdef HAVE_UCX_GPU_DEVICE_API +#include +#endif +} + +namespace nixl::ucx { + +#ifdef HAVE_UCX_GPU_DEVICE_API + +nixlGpuXferReqH +createGpuXferReq(const nixlUcxEp &ep, + const std::vector &local_mems, + const std::vector &remote_rkeys) { + nixl_status_t status = ep.checkTxState(); + if (status != NIXL_SUCCESS) { + throw std::runtime_error("Endpoint not in valid state for creating memory list"); + } + + if (local_mems.empty() || remote_rkeys.empty()) { + throw std::invalid_argument("Empty memh or rkey lists provided"); + } + + if (local_mems.size() != remote_rkeys.size()) { + throw std::invalid_argument("Local memh and remote rkey lists must have same size"); + } + + std::vector ucp_elements; + ucp_elements.reserve(local_mems.size()); + + for (size_t i = 0; i < local_mems.size(); i++) { + ucp_device_mem_list_elem_t ucp_elem; + ucp_elem.field_mask = + UCP_DEVICE_MEM_LIST_ELEM_FIELD_MEMH | UCP_DEVICE_MEM_LIST_ELEM_FIELD_RKEY; + ucp_elem.memh = local_mems[i].getMemh(); + ucp_elem.rkey = remote_rkeys[i]->get(); + ucp_elements.push_back(ucp_elem); + } + + ucp_device_mem_list_params_t params; + params.field_mask = UCP_DEVICE_MEM_LIST_PARAMS_FIELD_ELEMENTS | + UCP_DEVICE_MEM_LIST_PARAMS_FIELD_ELEMENT_SIZE | + UCP_DEVICE_MEM_LIST_PARAMS_FIELD_NUM_ELEMENTS; + params.elements = ucp_elements.data(); + params.element_size = sizeof(ucp_device_mem_list_elem_t); + params.num_elements = ucp_elements.size(); + + ucp_device_mem_list_handle_h ucx_handle; + ucs_status_t ucs_status = ucp_device_mem_list_create(ep.getEp(), ¶ms, &ucx_handle); + if (ucs_status != UCS_OK) { + throw std::runtime_error(std::string("Failed to create device memory list: ") + + ucs_status_string(ucs_status)); + } + + NIXL_DEBUG << "Created device memory list handle with " << local_mems.size() << " elements"; + return reinterpret_cast(ucx_handle); +} + +void +releaseGpuXferReq(nixlGpuXferReqH gpu_req) noexcept { + auto ucx_handle = reinterpret_cast(gpu_req); + ucp_device_mem_list_release(ucx_handle); +} + +#else + +nixlGpuXferReqH +createGpuXferReq(const nixlUcxEp &ep, + const std::vector &local_mems, + const std::vector &remote_rkeys) { + NIXL_ERROR << "UCX GPU device API not supported"; + throw std::runtime_error("UCX GPU device API not available"); +} + +void +releaseGpuXferReq(nixlGpuXferReqH gpu_req) noexcept { + NIXL_WARN << "UCX GPU device API not supported - cannot release GPU transfer request handle"; +} + +#endif + +} // namespace nixl::ucx diff --git a/src/utils/common/list_elem.h b/src/utils/ucx/gpu_xfer_req_h.h similarity index 53% rename from src/utils/common/list_elem.h rename to src/utils/ucx/gpu_xfer_req_h.h index a9f26bd2ff..11803e811f 100644 --- a/src/utils/common/list_elem.h +++ b/src/utils/ucx/gpu_xfer_req_h.h @@ -14,45 +14,27 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -#ifndef _NIXL_LIST_ELEM_H -#define _NIXL_LIST_ELEM_H - - -template -class nixlLinkElem { -private: - T *_next; -public: - nixlLinkElem() - { - _next = NULL; - } - - ~nixlLinkElem() { - _next = NULL;; - } - - /* Link this element into the chain after "elem" */ - void link(T *elem) - { - elem->_next = _next; - _next = elem; - } - - /* Exclude this element from the chain, return the new head */ - T *unlink() - { - T *ret = _next; - /* Forget my place */ - _next = NULL; - return ret; - } - - T *next() - { - return _next; - } - -} ; + +#ifndef NIXL_SRC_UTILS_UCX_GPU_XFER_REQ_H_H +#define NIXL_SRC_UTILS_UCX_GPU_XFER_REQ_H_H + +#include + +#include "nixl_types.h" + +class nixlUcxEp; +class nixlUcxMem; + +namespace nixl::ucx { +class rkey; + +nixlGpuXferReqH +createGpuXferReq(const nixlUcxEp &, + const std::vector &, + const std::vector &); + +void +releaseGpuXferReq(nixlGpuXferReqH gpu_req) noexcept; +} // namespace nixl::ucx #endif diff --git a/src/utils/ucx/meson.build b/src/utils/ucx/meson.build index 647701893c..cf70169441 100644 --- a/src/utils/ucx/meson.build +++ b/src/utils/ucx/meson.build @@ -21,8 +21,10 @@ endif ucx_utils_inc_dirs = include_directories('.') +ucx_utils_sources = ['ucx_utils.cpp', 'ucx_utils.h', 'config.h', 'config.cpp', 'rkey.cpp', 'rkey.h', 'gpu_xfer_req_h.cpp', 'gpu_xfer_req_h.h'] + ucx_utils_lib = library('ucx_utils', - 'ucx_utils.cpp', 'ucx_utils.h', 'config.h', 'config.cpp', 'rkey.cpp', 'rkey.h', + ucx_utils_sources, dependencies: ucx_utils_dep, include_directories: [ nixl_inc_dirs, utils_inc_dirs ], install: true) diff --git a/src/utils/ucx/ucx_utils.cpp b/src/utils/ucx/ucx_utils.cpp index 7f4dedb648..1bbac53c82 100644 --- a/src/utils/ucx/ucx_utils.cpp +++ b/src/utils/ucx/ucx_utils.cpp @@ -26,6 +26,12 @@ #include +extern "C" { +#ifdef HAVE_UCX_GPU_DEVICE_API +#include +#endif +} + #include "common/nixl_log.h" #include "config.h" #include "serdes/serdes.h" @@ -58,6 +64,8 @@ nixl_status_t ucx_status_to_nixl(ucs_status_t status) return NIXL_ERR_REMOTE_DISCONNECT; case UCS_ERR_INVALID_PARAM: return NIXL_ERR_INVALID_PARAM; + case UCS_ERR_CANCELED: + return NIXL_ERR_CANCELED; default: NIXL_WARN << "Unexpected UCX error: " << ucs_status_string(status); return NIXL_ERR_BACKEND; @@ -233,14 +241,17 @@ nixl_status_t nixlUcxEp::sendAm(unsigned msg_id, void* buffer, size_t len, uint32_t flags, nixlUcxReq &req) { - ucs_status_ptr_t request; + nixl_status_t status = checkTxState(); + if (status != NIXL_SUCCESS) { + return status; + } + ucp_request_param_t param = {0}; param.op_attr_mask |= UCP_OP_ATTR_FIELD_FLAGS; param.flags = flags; - request = ucp_am_send_nbx(eph, msg_id, hdr, hdr_len, buffer, len, ¶m); - + ucs_status_ptr_t request = ucp_am_send_nbx(eph, msg_id, hdr, hdr_len, buffer, len, ¶m); if (UCS_PTR_IS_PTR(request)) { req = (void*)request; return NIXL_IN_PROG; @@ -386,6 +397,10 @@ nixlUcxContext::nixlUcxContext(std::vector devs, ucp_params.field_mask = UCP_PARAM_FIELD_FEATURES | UCP_PARAM_FIELD_MT_WORKERS_SHARED; ucp_params.features = UCP_FEATURE_RMA | UCP_FEATURE_AMO32 | UCP_FEATURE_AMO64 | UCP_FEATURE_AM; +#ifdef HAVE_UCX_GPU_DEVICE_API + ucp_params.features |= UCP_FEATURE_DEVICE; +#endif + if (prog_thread) ucp_params.features |= UCP_FEATURE_WAKEUP; ucp_params.mt_workers_shared = num_workers > 1 ? 1 : 0; @@ -584,6 +599,34 @@ void nixlUcxContext::memDereg(nixlUcxMem &mem) ucp_mem_unmap(ctx, mem.memh); } +#ifndef HAVE_UCX_GPU_DEVICE_API +namespace { +constexpr std::string_view ucxGpuDeviceApiUnsupported{ + "UCX was not compiled with GPU device API support"}; +} +#endif + + + +size_t +nixlUcxContext::getGpuSignalSize() const { +#ifdef HAVE_UCX_GPU_DEVICE_API + ucp_context_attr_t attr; + attr.field_mask = UCP_ATTR_FIELD_DEVICE_COUNTER_SIZE; + ucs_status_t query_status = ucp_context_query(ctx, &attr); + + if (query_status != UCS_OK) { + throw std::runtime_error( + std::string("Failed to query UCX context for device counter size: ") + + ucs_status_string(query_status)); + } + + return attr.device_counter_size; +#else + throw std::runtime_error(std::string(ucxGpuDeviceApiUnsupported)); +#endif +} + /* =========================================== * Active message handling * =========================================== */ @@ -654,3 +697,26 @@ nixlUcxWorker::getEfd() const { } return fd; } + +void +nixlUcxWorker::prepGpuSignal([[maybe_unused]] const nixlUcxMem &mem, + [[maybe_unused]] void *signal) const { +#ifdef HAVE_UCX_GPU_DEVICE_API + if (!signal) { + throw std::invalid_argument("Signal pointer cannot be null"); + } + + ucp_device_counter_params_t params; + params.field_mask = UCP_DEVICE_COUNTER_PARAMS_FIELD_MEMH; + params.memh = mem.memh; + + ucs_status_t status = ucp_device_counter_init(worker.get(), ¶ms, signal); + + if (status != UCS_OK) { + throw std::runtime_error(std::string("Failed to initialize GPU signal: ") + + ucs_status_string(status)); + } +#else + throw std::runtime_error(std::string(ucxGpuDeviceApiUnsupported)); +#endif +} diff --git a/src/utils/ucx/ucx_utils.h b/src/utils/ucx/ucx_utils.h index 2a21cb661b..4cbe598b97 100644 --- a/src/utils/ucx/ucx_utils.h +++ b/src/utils/ucx/ucx_utils.h @@ -28,6 +28,8 @@ extern "C" #include #include "absl/status/statusor.h" +#include "absl/strings/numbers.h" + enum class nixl_ucx_mt_t { SINGLE, @@ -57,6 +59,24 @@ template return "INVALID"; // It is not a to_string function's job to validate. } +template +[[nodiscard]] T +nixl_b_params_get(const nixl_b_params_t *custom_params, const std::string &key, T default_value) { + if (!custom_params) { + return default_value; + } + + auto it = custom_params->find(key); + if (it == custom_params->end()) { + return default_value; + } + + if constexpr (std::is_same_v) { + T result; + return absl::SimpleAtoi(it->second, &result) ? result : default_value; + } +} + using nixlUcxReq = void*; namespace nixl::ucx { @@ -140,6 +160,11 @@ class nixlUcxMem { size_t size; ucp_mem_h memh; public: + [[nodiscard]] ucp_mem_h + getMemh() const noexcept { + return memh; + } + friend class nixlUcxWorker; friend class nixlUcxContext; friend class nixlUcxEp; @@ -167,6 +192,10 @@ class nixlUcxContext { [[nodiscard]] std::string packRkey(nixlUcxMem &mem); void memDereg(nixlUcxMem &mem); + /* GPU signal management */ + [[nodiscard]] size_t + getGpuSignalSize() const; + friend class nixlUcxWorker; }; @@ -203,6 +232,10 @@ class nixlUcxWorker { [[nodiscard]] int getEfd() const; + /* GPU signal management */ + void + prepGpuSignal(const nixlUcxMem &mem, void *signal) const; + private: [[nodiscard]] static ucp_worker * createUcpWorker(const nixlUcxContext &); diff --git a/test/README.md b/test/README.md index 1da5eb5958..6abe583b7c 100644 --- a/test/README.md +++ b/test/README.md @@ -10,6 +10,33 @@ Here are all the explained tests in this directory. There are more specific unit - test/ucx_backend_multi.cpp - Multi threaded test of UCX connection setup/teardown - test/python/nixl_bindings_test.py - single threaded Python test of nixlAgent, nixlBasicDesc, and nixlDescList python bindings +## Google Test Framework (gtest) + +The project includes comprehensive unit tests using Google Test framework located in `test/gtest/`: + +- test/gtest/telemetry_test.cpp - Comprehensive tests for NIXL telemetry functionality including initialization, data tracking, thread safety, and edge cases +- test/gtest/query_mem.cpp - Tests for memory query functionality +- test/gtest/error_handling.cpp - Tests for error handling and status codes +- test/gtest/test_transfer.cpp - Tests for data transfer operations +- test/gtest/plugin_manager.cpp - Tests for plugin management +- test/gtest/multi_threading.cpp - Multi-threaded test scenarios +- test/gtest/metadata_exchange.cpp - Tests for metadata exchange functionality + +To run the gtest suite: +```bash +# Build the project +meson setup build +meson compile -C build + +# Run all tests +cd build +./gtest + +# Run specific test categories +./gtest --gtest_filter="TelemetryTest*" # Run only telemetry tests +./gtest --gtest_filter="QueryMemTest*" # Run only query memory tests +``` + # NIXL_wrapper python class To make the NIXL interface more python style, a wrapper class was added on top of python bindings. diff --git a/test/gtest/common.cpp b/test/gtest/common.cpp index 50f6d1ba6e..b763f5ee48 100644 --- a/test/gtest/common.cpp +++ b/test/gtest/common.cpp @@ -19,8 +19,14 @@ #include #include #include +#include +#include #include #include +#include +#include +#include +#include namespace gtest { @@ -39,6 +45,11 @@ void ScopedEnv::addVar(const std::string &name, const std::string &value) m_vars.emplace(name, value); } +void +ScopedEnv::popVar() { + m_vars.pop(); +} + ScopedEnv::Variable::Variable(const std::string &name, const std::string &value) : m_name(name) { @@ -72,4 +83,54 @@ ScopedEnv::Variable::~Variable() } } +PortAllocator & +PortAllocator::instance() { + static PortAllocator _instance; + return _instance; +} + +void +PortAllocator::set_min_port(uint16_t min_port) { + _min_port = min_port; + _port = _min_port; +} + +void +PortAllocator::set_max_port(uint16_t max_port) { + _max_port = max_port; +} + +bool +PortAllocator::is_port_available(uint16_t port) { + struct sockaddr_in addr = { + .sin_family = AF_INET, .sin_port = htons(port), .sin_addr = {.s_addr = INADDR_ANY}}; + + const auto sock_fd = socket(AF_INET, SOCK_STREAM, 0); + const auto ret = bind(sock_fd, (struct sockaddr *)&addr, sizeof(addr)); + close(sock_fd); + return ret == 0; +} + +uint16_t +PortAllocator::next_tcp_port() { + PortAllocator &instance = PortAllocator::instance(); + std::lock_guard lock(instance._mutex); + const int port_range = instance._max_port - instance._min_port; + + for (int scanned = 0; scanned < port_range; scanned++) { + if (is_port_available(instance._port)) { + return instance._port++; + } + + instance._port++; + + if (instance._port >= instance._max_port) { + instance._port = instance._min_port; + } + } + + throw std::runtime_error("No port available in range: " + std::to_string(instance._min_port) + + " - " + std::to_string(instance._max_port)); +} + } // namespace gtest diff --git a/test/gtest/common.h b/test/gtest/common.h index 2355a748f5..ef0a1cf2bd 100644 --- a/test/gtest/common.h +++ b/test/gtest/common.h @@ -20,8 +20,11 @@ #include #include #include +#include +#include #include #include +#include namespace gtest { constexpr const char * @@ -43,7 +46,10 @@ class Logger { class ScopedEnv { public: - void addVar(const std::string &name, const std::string &value); + void + addVar(const std::string &name, const std::string &value); + void + popVar(); private: class Variable { @@ -63,6 +69,39 @@ class ScopedEnv { std::stack m_vars; }; +class PortAllocator { +public: + static constexpr uint16_t MIN_PORT = 10500; + static constexpr uint16_t MAX_PORT = 65535; + +private: + PortAllocator() = default; + ~PortAllocator() = default; + PortAllocator(const PortAllocator &other) = delete; + void + operator=(const PortAllocator &) = delete; + +public: + static uint16_t + next_tcp_port(); + static PortAllocator & + instance(); + + void + set_min_port(uint16_t min_port); + void + set_max_port(uint16_t max_port); + +private: + static bool + is_port_available(uint16_t port); + + std::mutex _mutex; + uint16_t _port = MIN_PORT; + uint16_t _min_port = MIN_PORT; + uint16_t _max_port = MAX_PORT; +}; + } // namespace gtest #endif /* TEST_GTEST_COMMON_H */ diff --git a/test/gtest/device_api/meson.build b/test/gtest/device_api/meson.build new file mode 100644 index 0000000000..0fa19e0f37 --- /dev/null +++ b/test/gtest/device_api/meson.build @@ -0,0 +1,25 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Device API test utilities and sources +device_api_inc_dirs = include_directories('.') + +device_api_sources = [ + 'single_write_test.cu', + 'utils.cu' +] + +# Export for parent meson.build +device_api_test_sources = files(device_api_sources) diff --git a/test/gtest/device_api/single_write_test.cu b/test/gtest/device_api/single_write_test.cu new file mode 100644 index 0000000000..086d1547f3 --- /dev/null +++ b/test/gtest/device_api/single_write_test.cu @@ -0,0 +1,585 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "utils.cuh" + +namespace gtest::nixl::gpu::single_write { + +template +__global__ void +TestSingleWriteKernel(nixlGpuXferReqH req_hdnl, + unsigned index, + const void *src_addr, + uint64_t remote_addr, + size_t size, + size_t num_iters, + bool is_no_delay, + unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr) { + __shared__ nixlGpuXferStatusH xfer_status[MAX_THREADS]; + nixlGpuXferStatusH *xfer_status_ptr = &xfer_status[GetReqIdx()]; + nixl_status_t status; + + assert(GetReqIdx() < MAX_THREADS); + + if (threadIdx.x == 0) { + unsigned long long start_time = GetTimeNs(); + *start_time_ptr = start_time; + } + + __syncthreads(); + + for (size_t i = 0; i < num_iters; ++i) { + status = nixlGpuPostSingleWriteXferReq( + req_hdnl, index, src_addr, remote_addr, size, is_no_delay, xfer_status_ptr); + if (status != NIXL_SUCCESS) { + printf("Thread %d: nixlGpuPostSingleWriteXferReq failed iteration %lu: status=%d (0x%x)\n", + threadIdx.x, + (unsigned long)i, + status, + static_cast(status)); + return; + } + + status = nixlGpuGetXferStatus(*xfer_status_ptr); + if (status != NIXL_SUCCESS && status != NIXL_IN_PROG) { + printf("Thread %d: Failed to progress single write transfer iteration %zu: status=%d\n", + threadIdx.x, + i, + status); + return; + } + + while (status == NIXL_IN_PROG) { + status = nixlGpuGetXferStatus(*xfer_status_ptr); + if (status != NIXL_SUCCESS && status != NIXL_IN_PROG) { + printf("Thread %d: Failed to progress single write transfer iteration %zu: status=%d\n", + threadIdx.x, + i, + status); + return; + } + } + + if (status != NIXL_SUCCESS) { + printf("Thread %d: Transfer completion failed iteration %zu: status=%d\n", + threadIdx.x, + i, + status); + return; + } + } + + if (threadIdx.x == 0) { + unsigned long long end_time = GetTimeNs(); + *end_time_ptr = end_time; + } +} + +template +nixl_status_t +LaunchSingleWriteTest(unsigned num_threads, + nixlGpuXferReqH req_hdnl, + unsigned index, + const void *src_addr, + uint64_t remote_addr, + size_t size, + size_t num_iters, + bool is_no_delay, + unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr) { + nixl_status_t ret = NIXL_SUCCESS; + cudaError_t err; + + TestSingleWriteKernel<<<1, num_threads>>>(req_hdnl, + index, + src_addr, + remote_addr, + size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + + err = cudaDeviceSynchronize(); + if (err != cudaSuccess) { + printf("Failed to synchronize: %s\n", cudaGetErrorString(err)); + ret = NIXL_ERR_BACKEND; + } + + err = cudaGetLastError(); + if (err != cudaSuccess) { + printf("Failed to launch kernel: %s\n", cudaGetErrorString(err)); + ret = NIXL_ERR_BACKEND; + } + + return ret; +} + +class SingleWriteTest : public DeviceApiTestBase { +protected: + std::string getBackendName() const { return "UCX"; } + + static nixlAgentConfig + getConfig() { + return nixlAgentConfig(true, + false, + 0, + nixl_thread_sync_t::NIXL_THREAD_SYNC_RW, + 0, + 100000); + } + + nixl_b_params_t + getBackendParams() { + nixl_b_params_t params; + + if (getBackendName() == "UCX") { + params["num_workers"] = "2"; + } + + return params; + } + + void + SetUp() override { + if (cudaSetDevice(0) != cudaSuccess) { + FAIL() << "Failed to set CUDA device 0"; + } + + for (size_t i = 0; i < 2; i++) { + agents.emplace_back(std::make_unique(getAgentName(i), getConfig())); + nixlBackendH *backend_handle = nullptr; + nixl_status_t status = + agents.back()->createBackend(getBackendName(), getBackendParams(), backend_handle); + ASSERT_EQ(status, NIXL_SUCCESS); + EXPECT_NE(backend_handle, nullptr); + backend_handles.push_back(backend_handle); + } + } + + void + TearDown() override { + agents.clear(); + backend_handles.clear(); + } + + template + nixlDescList + makeDescList(const std::vector &buffers, nixl_mem_t mem_type) { + nixlDescList desc_list(mem_type); + for (const auto &buffer : buffers) { + desc_list.addDesc(Desc(buffer, buffer.getSize(), uint64_t(DEV_ID))); + } + return desc_list; + } + + void + registerMem(nixlAgent &agent, const std::vector &buffers, nixl_mem_t mem_type) { + auto reg_list = makeDescList(buffers, mem_type); + agent.registerMem(reg_list); + } + + void + completeWireup(size_t from_agent, size_t to_agent) { + nixl_notifs_t notifs; + nixl_status_t status = getAgent(from_agent).genNotif(getAgentName(to_agent), NOTIF_MSG); + ASSERT_EQ(status, NIXL_SUCCESS) << "Failed to complete wireup"; + + do { + nixl_status_t ret = getAgent(to_agent).getNotifs(notifs); + ASSERT_EQ(ret, NIXL_SUCCESS) << "Failed to get notifications during wireup"; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } while (notifs.size() == 0); + } + + void + exchangeMD(size_t from_agent, size_t to_agent) { + for (size_t i = 0; i < agents.size(); i++) { + nixl_blob_t md; + nixl_status_t status = agents[i]->getLocalMD(md); + ASSERT_EQ(status, NIXL_SUCCESS); + + for (size_t j = 0; j < agents.size(); j++) { + if (i == j) continue; + std::string remote_agent_name; + status = agents[j]->loadRemoteMD(md, remote_agent_name); + ASSERT_EQ(status, NIXL_SUCCESS); + EXPECT_EQ(remote_agent_name, getAgentName(i)); + } + } + + completeWireup(from_agent, to_agent); + } + + void + invalidateMD() { + for (size_t i = 0; i < agents.size(); i++) { + for (size_t j = 0; j < agents.size(); j++) { + if (i == j) continue; + nixl_status_t status = agents[j]->invalidateRemoteMD(getAgentName(i)); + ASSERT_EQ(status, NIXL_SUCCESS); + } + } + } + + void + createRegisteredMem(nixlAgent &agent, + size_t size, + size_t count, + nixl_mem_t mem_type, + std::vector &out) { + while (count-- != 0) { + out.emplace_back(size, mem_type); + } + + registerMem(agent, out, mem_type); + } + + nixlAgent & + getAgent(size_t idx) { + return *agents[idx]; + } + + std::string + getAgentName(size_t idx) { + return absl::StrFormat("agent_%d", idx); + } + + nixl_status_t + dispatchLaunchSingleWriteTest(nixl_gpu_level_t level, + unsigned num_threads, + nixlGpuXferReqH req_hdnl, + unsigned index, + const void *src_addr, + uint64_t remote_addr, + size_t size, + size_t num_iters, + bool is_no_delay, + unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr) { + switch (level) { + case nixl_gpu_level_t::BLOCK: + return LaunchSingleWriteTest(num_threads, + req_hdnl, + index, + src_addr, + remote_addr, + size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + case nixl_gpu_level_t::WARP: + return LaunchSingleWriteTest(num_threads, + req_hdnl, + index, + src_addr, + remote_addr, + size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + case nixl_gpu_level_t::THREAD: + return LaunchSingleWriteTest( + num_threads, + req_hdnl, + index, + src_addr, + remote_addr, + size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + default: + ADD_FAILURE() << "Unknown level: " << static_cast(level); + return NIXL_ERR_INVALID_PARAM; + } + } + +protected: + static constexpr size_t SENDER_AGENT = 0; + static constexpr size_t RECEIVER_AGENT = 1; + +private: + static constexpr uint64_t DEV_ID = 0; + + std::vector> agents; + std::vector backend_handles; + + void + initTiming(unsigned long long **start_time_ptr, unsigned long long **end_time_ptr) { + cudaMalloc(start_time_ptr, sizeof(unsigned long long)); + cudaMalloc(end_time_ptr, sizeof(unsigned long long)); + cudaMemset(*start_time_ptr, 0, sizeof(unsigned long long)); + cudaMemset(*end_time_ptr, 0, sizeof(unsigned long long)); + } + + void + getTiming(unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr, + unsigned long long &start_time_cpu, + unsigned long long &end_time_cpu) { + cudaMemcpy( + &start_time_cpu, start_time_ptr, sizeof(unsigned long long), cudaMemcpyDeviceToHost); + cudaMemcpy(&end_time_cpu, end_time_ptr, sizeof(unsigned long long), cudaMemcpyDeviceToHost); + } + + void + logResults(size_t size, + size_t count, + size_t num_iters, + unsigned long long start_time_cpu, + unsigned long long end_time_cpu) { + auto total_time = NS_TO_SEC(end_time_cpu - start_time_cpu); + double total_size = size * count * num_iters; + auto bandwidth = total_size / total_time / (1024 * 1024); + Logger() << "SingleWrite Results: " << size << "x" << count << "x" << num_iters << "=" + << total_size << " bytes in " << total_time << " seconds " << "(" << bandwidth + << " MB/s)"; + } + +public: + void + initTimingPublic(unsigned long long **start_time_ptr, unsigned long long **end_time_ptr) { + initTiming(start_time_ptr, end_time_ptr); + } + + void + getTimingPublic(unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr, + unsigned long long &start_time_cpu, + unsigned long long &end_time_cpu) { + getTiming(start_time_ptr, end_time_ptr, start_time_cpu, end_time_cpu); + } + + void + logResultsPublic(size_t size, + size_t count, + size_t num_iters, + unsigned long long start_time_cpu, + unsigned long long end_time_cpu) { + logResults(size, count, num_iters, start_time_cpu, end_time_cpu); + } +}; + +TEST_P(SingleWriteTest, BasicSingleWriteTest) { + std::vector src_buffers, dst_buffers; + constexpr size_t size = 4 * 1024; + constexpr size_t count = 1; + nixl_mem_t mem_type = VRAM_SEG; + size_t num_threads = 32; + const size_t num_iters = 10000; + constexpr unsigned index = 0; + const bool is_no_delay = true; + + createRegisteredMem(getAgent(SENDER_AGENT), size, count, mem_type, src_buffers); + createRegisteredMem(getAgent(RECEIVER_AGENT), size, count, mem_type, dst_buffers); + + uint32_t *src_data = static_cast(static_cast(src_buffers[0])); + uint32_t pattern = 0xDEADBEEF; + + cudaMemset(src_data, 0, size); + cudaMemcpy(src_data, &pattern, sizeof(pattern), cudaMemcpyHostToDevice); + + exchangeMD(SENDER_AGENT, RECEIVER_AGENT); + + nixl_opt_args_t extra_params = {}; + extra_params.hasNotif = true; + extra_params.notifMsg = NOTIF_MSG; + + nixlXferReqH *xfer_req = nullptr; + nixl_status_t status = getAgent(SENDER_AGENT) + .createXferReq(NIXL_WRITE, + makeDescList(src_buffers, mem_type), + makeDescList(dst_buffers, mem_type), + getAgentName(RECEIVER_AGENT), + xfer_req, + &extra_params); + + ASSERT_EQ(status, NIXL_SUCCESS) + << "Failed to create xfer request " << nixlEnumStrings::statusStr(status); + EXPECT_NE(xfer_req, nullptr); + + nixlGpuXferReqH gpu_req_hndl; + status = getAgent(SENDER_AGENT).createGpuXferReq(*xfer_req, gpu_req_hndl); + ASSERT_EQ(status, NIXL_SUCCESS) << "Failed to create GPU xfer request"; + + ASSERT_NE(gpu_req_hndl, nullptr) << "GPU request handle is null after createGpuXferReq"; + + uint64_t remote_addr = static_cast(dst_buffers[0]); + const void *src_addr = static_cast(src_buffers[0]); + + unsigned long long *start_time_ptr = nullptr; + unsigned long long *end_time_ptr = nullptr; + nixl_status_t *result_status = nullptr; + + initTimingPublic(&start_time_ptr, &end_time_ptr); + cudaMalloc(&result_status, sizeof(nixl_status_t)); + cudaMemset(result_status, 0, sizeof(nixl_status_t)); + + status = dispatchLaunchSingleWriteTest(GetParam(), + num_threads, + gpu_req_hndl, + index, + src_addr, + remote_addr, + size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + + ASSERT_EQ(status, NIXL_SUCCESS) << "Kernel launch failed with status: " << status; + + nixl_status_t gpu_result; + cudaMemcpy(&gpu_result, result_status, sizeof(nixl_status_t), cudaMemcpyDeviceToHost); + ASSERT_EQ(gpu_result, NIXL_SUCCESS) << "GPU kernel reported error: " << gpu_result; + + unsigned long long start_time_cpu = 0; + unsigned long long end_time_cpu = 0; + getTimingPublic(start_time_ptr, end_time_ptr, start_time_cpu, end_time_cpu); + logResultsPublic(size, count, num_iters, start_time_cpu, end_time_cpu); + + uint32_t dst_data; + cudaMemcpy(&dst_data, + static_cast(static_cast(dst_buffers[0])), + sizeof(uint32_t), + cudaMemcpyDeviceToHost); + EXPECT_EQ(dst_data, pattern) << "Data transfer verification failed. Expected: 0x" << std::hex + << pattern << ", Got: 0x" << dst_data; + + cudaFree(start_time_ptr); + cudaFree(end_time_ptr); + cudaFree(result_status); + + getAgent(SENDER_AGENT).releaseGpuXferReq(gpu_req_hndl); + + status = getAgent(SENDER_AGENT).releaseXferReq(xfer_req); + EXPECT_EQ(status, NIXL_SUCCESS); + + invalidateMD(); +} + +TEST_P(SingleWriteTest, VariableSizeTest) { + std::vector test_sizes = {64, 256, 1024, 4096, 16384}; + + for (size_t test_size : test_sizes) { + std::vector src_buffers, dst_buffers; + constexpr size_t count = 1; + nixl_mem_t mem_type = VRAM_SEG; + size_t num_threads = 32; + const size_t num_iters = 50000; + constexpr unsigned index = 0; + const bool is_no_delay = true; + + createRegisteredMem(getAgent(SENDER_AGENT), test_size, count, mem_type, src_buffers); + createRegisteredMem(getAgent(RECEIVER_AGENT), test_size, count, mem_type, dst_buffers); + + std::vector pattern(test_size); + for (size_t i = 0; i < test_size; ++i) { + pattern[i] = static_cast(i % 256); + } + + cudaMemcpy( + static_cast(src_buffers[0]), pattern.data(), test_size, cudaMemcpyHostToDevice); + + exchangeMD(SENDER_AGENT, RECEIVER_AGENT); + + nixl_opt_args_t extra_params = {}; + extra_params.hasNotif = true; + extra_params.notifMsg = NOTIF_MSG; + + nixlXferReqH *xfer_req = nullptr; + nixl_status_t status = + getAgent(SENDER_AGENT) + .createXferReq(NIXL_WRITE, + makeDescList(src_buffers, mem_type), + makeDescList(dst_buffers, mem_type), + getAgentName(RECEIVER_AGENT), + xfer_req, + &extra_params); + + ASSERT_EQ(status, NIXL_SUCCESS) << "Failed to create xfer request for size " << test_size; + + nixlGpuXferReqH gpu_req_hndl; + status = getAgent(SENDER_AGENT).createGpuXferReq(*xfer_req, gpu_req_hndl); + ASSERT_EQ(status, NIXL_SUCCESS) + << "Failed to create GPU xfer request for size " << test_size; + + ASSERT_NE(gpu_req_hndl, nullptr) << "GPU request handle is null after createGpuXferReq"; + + unsigned long long *start_time_ptr = nullptr; + unsigned long long *end_time_ptr = nullptr; + nixl_status_t *result_status = nullptr; + + initTimingPublic(&start_time_ptr, &end_time_ptr); + cudaMalloc(&result_status, sizeof(nixl_status_t)); + cudaMemset(result_status, 0, sizeof(nixl_status_t)); + + uint64_t remote_addr = static_cast(dst_buffers[0]); + const void *src_addr = static_cast(src_buffers[0]); + + status = dispatchLaunchSingleWriteTest(GetParam(), + num_threads, + gpu_req_hndl, + index, + src_addr, + remote_addr, + test_size, + num_iters, + is_no_delay, + start_time_ptr, + end_time_ptr); + + ASSERT_EQ(status, NIXL_SUCCESS) << "Kernel launch failed for size " << test_size; + + nixl_status_t gpu_result; + cudaMemcpy(&gpu_result, result_status, sizeof(nixl_status_t), cudaMemcpyDeviceToHost); + ASSERT_EQ(gpu_result, NIXL_SUCCESS) << "GPU kernel failed for size " << test_size; + + std::vector received_data(test_size); + cudaMemcpy(received_data.data(), + static_cast(dst_buffers[0]), + test_size, + cudaMemcpyDeviceToHost); + + EXPECT_EQ(received_data, pattern) << "Data verification failed for size " << test_size; + + cudaFree(start_time_ptr); + cudaFree(end_time_ptr); + cudaFree(result_status); + + getAgent(SENDER_AGENT).releaseGpuXferReq(gpu_req_hndl); + getAgent(SENDER_AGENT).releaseXferReq(xfer_req); + invalidateMD(); + } +} + +} // namespace gtest::nixl::gpu::single_write + +using gtest::nixl::gpu::single_write::SingleWriteTest; + +INSTANTIATE_TEST_SUITE_P( + ucxDeviceApi, + SingleWriteTest, + testing::ValuesIn(gtest::gpu::_test_levels), + [](const testing::TestParamInfo &info) { + return std::string("UCX_") + gtest::gpu::GetGpuXferLevelStr(info.param); + }); diff --git a/test/gtest/device_api/utils.cu b/test/gtest/device_api/utils.cu new file mode 100644 index 0000000000..acbc603133 --- /dev/null +++ b/test/gtest/device_api/utils.cu @@ -0,0 +1,163 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "utils.cuh" + +namespace gtest { +namespace gpu { + +const char *GetGpuXferLevelStr(nixl_gpu_level_t level) { + switch (level) { + case nixl_gpu_level_t::WARP: + return "WARP"; + case nixl_gpu_level_t::BLOCK: + return "BLOCK"; + case nixl_gpu_level_t::THREAD: + return "THREAD"; + default: + return "UNKNOWN"; + } +} + +void initTiming(unsigned long long **start_time_ptr, unsigned long long **end_time_ptr) { + cudaMalloc(start_time_ptr, sizeof(unsigned long long)); + cudaMalloc(end_time_ptr, sizeof(unsigned long long)); + cudaMemset(*start_time_ptr, 0, sizeof(unsigned long long)); + cudaMemset(*end_time_ptr, 0, sizeof(unsigned long long)); +} + +void getTiming(unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr, + unsigned long long &start_time_cpu, + unsigned long long &end_time_cpu) { + cudaMemcpy(&start_time_cpu, start_time_ptr, sizeof(unsigned long long), cudaMemcpyDeviceToHost); + cudaMemcpy(&end_time_cpu, end_time_ptr, sizeof(unsigned long long), cudaMemcpyDeviceToHost); +} + +void logResults(size_t size, + size_t count, + size_t num_iters, + unsigned long long start_time_cpu, + unsigned long long end_time_cpu) { + auto total_time = NS_TO_SEC(end_time_cpu - start_time_cpu); + double total_size = size * count * num_iters; + auto bandwidth = total_size / total_time / (1024 * 1024); + printf("Device API Results: %zux%zux%zu=%.0f bytes in %f seconds (%.2f MB/s)\n", + size, count, num_iters, total_size, total_time, bandwidth); +} + +} // namespace gpu +} // namespace gtest + +nixlAgentConfig DeviceApiTestBase::getConfig() { + return nixlAgentConfig(true, + false, + 0, + nixl_thread_sync_t::NIXL_THREAD_SYNC_RW, + 0, + 100000); +} + +nixl_b_params_t DeviceApiTestBase::getBackendParams() { + nixl_b_params_t params; + params["num_workers"] = "2"; + return params; +} + +void DeviceApiTestBase::SetUp() { + if (cudaSetDevice(0) != cudaSuccess) { + FAIL() << "Failed to set CUDA device 0"; + } + + for (size_t i = 0; i < 2; i++) { + agents.emplace_back(std::make_unique(getAgentName(i), getConfig())); + nixlBackendH *backend_handle = nullptr; + nixl_status_t status = agents.back()->createBackend("UCX", getBackendParams(), backend_handle); + ASSERT_EQ(status, NIXL_SUCCESS); + EXPECT_NE(backend_handle, nullptr); + backend_handles.push_back(backend_handle); + } +} + +void DeviceApiTestBase::TearDown() { + agents.clear(); + backend_handles.clear(); +} + +template +nixlDescList DeviceApiTestBase::makeDescList(const std::vector &buffers, nixl_mem_t mem_type) { + nixlDescList desc_list(mem_type); + for (const auto &buffer : buffers) { + desc_list.addDesc(Desc(buffer, buffer.getSize(), uint64_t(DEV_ID))); + } + return desc_list; +} + +void DeviceApiTestBase::registerMem(nixlAgent &agent, const std::vector &buffers, nixl_mem_t mem_type) { + auto reg_list = makeDescList(buffers, mem_type); + agent.registerMem(reg_list); +} + +void DeviceApiTestBase::completeWireup(size_t from_agent, size_t to_agent) { + nixl_notifs_t notifs; + nixl_status_t status = getAgent(from_agent).genNotif(getAgentName(to_agent), NOTIF_MSG); + ASSERT_EQ(status, NIXL_SUCCESS) << "Failed to complete wireup"; + + do { + nixl_status_t ret = getAgent(to_agent).getNotifs(notifs); + ASSERT_EQ(ret, NIXL_SUCCESS) << "Failed to get notifications during wireup"; + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } while (notifs.size() == 0); +} + +void DeviceApiTestBase::exchangeMD(size_t from_agent, size_t to_agent) { + for (size_t i = 0; i < agents.size(); i++) { + nixl_blob_t md; + nixl_status_t status = agents[i]->getLocalMD(md); + ASSERT_EQ(status, NIXL_SUCCESS); + + for (size_t j = 0; j < agents.size(); j++) { + if (i == j) continue; + std::string remote_agent_name; + status = agents[j]->loadRemoteMD(md, remote_agent_name); + ASSERT_EQ(status, NIXL_SUCCESS); + EXPECT_EQ(remote_agent_name, getAgentName(i)); + } + } + + completeWireup(from_agent, to_agent); +} + +void DeviceApiTestBase::invalidateMD() { + for (size_t i = 0; i < agents.size(); i++) { + for (size_t j = 0; j < agents.size(); j++) { + if (i == j) continue; + nixl_status_t status = agents[j]->invalidateRemoteMD(getAgentName(i)); + ASSERT_EQ(status, NIXL_SUCCESS); + } + } +} + +void DeviceApiTestBase::createRegisteredMem(nixlAgent &agent, + size_t size, + size_t count, + nixl_mem_t mem_type, + std::vector &out) { + while (count-- != 0) { + out.emplace_back(size, mem_type); + } + + registerMem(agent, out, mem_type); +} + +nixlAgent &DeviceApiTestBase::getAgent(size_t idx) { + return *agents[idx]; +} + +std::string DeviceApiTestBase::getAgentName(size_t idx) { + return absl::StrFormat("agent_%d", idx); +} + +template nixlDescList DeviceApiTestBase::makeDescList(const std::vector &buffers, nixl_mem_t mem_type); diff --git a/test/gtest/device_api/utils.cuh b/test/gtest/device_api/utils.cuh new file mode 100644 index 0000000000..9f3489980d --- /dev/null +++ b/test/gtest/device_api/utils.cuh @@ -0,0 +1,166 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef _DEVICE_API_UTILS_CUH +#define _DEVICE_API_UTILS_CUH + +#include +#include "nixl.h" +#include "common.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define MAX_THREADS 1024 +#define UCS_NSEC_PER_SEC 1000000000ul +#define NS_TO_SEC(ns) ((ns) * 1.0 / (UCS_NSEC_PER_SEC)) + +static const std::string NOTIF_MSG = "notification"; + +class MemBuffer : std::shared_ptr { +public: + MemBuffer(size_t size, nixl_mem_t mem_type) : + std::shared_ptr(allocate(size, mem_type), + [mem_type](void *ptr) { + release(ptr, mem_type); + }), + size(size) + { + } + + operator uintptr_t() const + { + return reinterpret_cast(get()); + } + + operator void*() const + { + return get(); + } + + operator const void*() const + { + return get(); + } + + size_t getSize() const + { + return size; + } + +private: + static void *allocate(size_t size, nixl_mem_t mem_type) + { + void *ptr; + return cudaSuccess == cudaMalloc(&ptr, size)? ptr : nullptr; + } + + static void release(void *ptr, nixl_mem_t mem_type) + { + cudaFree(ptr); + } + + size_t size; +}; + +namespace gtest { +namespace gpu { + +static const std::vector _test_levels = { + nixl_gpu_level_t::BLOCK, + nixl_gpu_level_t::WARP, + nixl_gpu_level_t::THREAD, +}; + +const char *GetGpuXferLevelStr(nixl_gpu_level_t level); + +void initTiming(unsigned long long **start_time_ptr, unsigned long long **end_time_ptr); +void getTiming(unsigned long long *start_time_ptr, + unsigned long long *end_time_ptr, + unsigned long long &start_time_cpu, + unsigned long long &end_time_cpu); +void logResults(size_t size, + size_t count, + size_t num_iters, + unsigned long long start_time_cpu, + unsigned long long end_time_cpu); + +} // namespace gpu +} // namespace gtest + +__device__ inline unsigned long long GetTimeNs() { + unsigned long long globaltimer; + asm volatile("mov.u64 %0, %globaltimer;" : "=l"(globaltimer)); + return globaltimer; +} + +template +__device__ constexpr size_t GetReqIdx() { + switch (level) { + case nixl_gpu_level_t::THREAD: + return threadIdx.x; + case nixl_gpu_level_t::WARP: + return threadIdx.x / warpSize; + case nixl_gpu_level_t::BLOCK: + return 0; + default: + return 0; + } +} + +class DeviceApiTestBase : public testing::TestWithParam { +protected: + static nixlAgentConfig getConfig(); + nixl_b_params_t getBackendParams(); + void SetUp() override; + void TearDown() override; + + template + nixlDescList makeDescList(const std::vector &buffers, nixl_mem_t mem_type); + + void registerMem(nixlAgent &agent, const std::vector &buffers, nixl_mem_t mem_type); + void completeWireup(size_t from_agent, size_t to_agent); + void exchangeMD(size_t from_agent, size_t to_agent); + void invalidateMD(); + + void createRegisteredMem(nixlAgent &agent, + size_t size, + size_t count, + nixl_mem_t mem_type, + std::vector &out); + + nixlAgent &getAgent(size_t idx); + std::string getAgentName(size_t idx); + +protected: + static constexpr size_t SENDER_AGENT = 0; + static constexpr size_t RECEIVER_AGENT = 1; + +private: + static constexpr uint64_t DEV_ID = 0; + + std::vector> agents; + std::vector backend_handles; +}; + +#endif // _DEVICE_API_UTILS_CUH diff --git a/test/gtest/device_api_test.cu b/test/gtest/device_api_test.cu new file mode 100644 index 0000000000..7bb16a437d --- /dev/null +++ b/test/gtest/device_api_test.cu @@ -0,0 +1,49 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include +#include "nixl.h" + +#include +#include + +namespace gtest::nixl::gpu { + +class ucxDeviceApi : public ::testing::Test {}; + +template +__global__ void dummyKernel() { + const nixlGpuSignal signal { 1, 0x1000 }; + nixlGpuXferStatusH status; + + [[maybe_unused]] auto result1 = nixlGpuPostSingleWriteXferReq(nullptr, 0, nullptr, 0, 0); + [[maybe_unused]] auto result2 = nixlGpuPostSignalXferReq(nullptr, 0, signal); + [[maybe_unused]] auto result3 = nixlGpuPostPartialWriteXferReq(nullptr, 0, nullptr, nullptr, nullptr, nullptr, signal); + [[maybe_unused]] auto result4 = nixlGpuPostWriteXferReq(nullptr, nullptr, nullptr, nullptr, signal); + [[maybe_unused]] auto result5 = nixlGpuGetXferStatus(status); + [[maybe_unused]] auto result6 = nixlGpuReadSignal(nullptr); +} + +TEST_F(ucxDeviceApi, compilationTest) { + dummyKernel<<<1, 1>>>(); + cudaError_t err = cudaGetLastError(); + ASSERT_EQ(err, cudaSuccess) << "Kernel launch failed: " << cudaGetErrorString(err); + + ASSERT_EQ(cudaDeviceSynchronize(), cudaSuccess); +} + +} // namespace gtest::nixl::gpu diff --git a/test/gtest/error_handling.cpp b/test/gtest/error_handling.cpp index 8fd500d640..2d37542875 100644 --- a/test/gtest/error_handling.cpp +++ b/test/gtest/error_handling.cpp @@ -26,8 +26,11 @@ namespace nixl { constexpr const char* ucx_err_handling_mode_key = "ucx_error_handling_mode"; constexpr const char* ucx_err_handling_mode_peer = "peer"; - static nixlBackendH* createUcxBackend(nixlAgent& agent, const std::string& backend_name) - { + static nixlBackendH * + createUcxBackend(nixlAgent &agent, + const std::string &backend_name, + size_t num_workers, + size_t num_threads) { std::vector plugins; nixl_status_t status = agent.getAvailPlugins(plugins); EXPECT_EQ(status, NIXL_SUCCESS); @@ -41,6 +44,10 @@ namespace nixl { nixlBackendH* backend_handle = nullptr; EXPECT_EQ(ucx_err_handling_mode_peer, params[ucx_err_handling_mode_key]); + params["num_workers"] = std::to_string(num_workers); + params["num_threads"] = std::to_string(num_threads); + // If threadpool is configured always force split + params["split_batch_size"] = "0"; status = agent.createBackend(*it, params, backend_handle); EXPECT_EQ(NIXL_SUCCESS, status); EXPECT_NE(nullptr, backend_handle); @@ -57,7 +64,8 @@ namespace nixl { } } // namespace nixl -class TestErrorHandling : public testing::TestWithParam { +// Tuple fields are: backend_name, num_workers, num_threads +class TestErrorHandling : public testing::TestWithParam> { class Agent { struct MemDesc { MemDesc() : m_dlist(DRAM_SEG), m_desc() {} @@ -79,8 +87,14 @@ class TestErrorHandling : public testing::TestWithParam { }; public: - void init(const std::string& name, const std::string& backend_name); - void destroy(); + void + init(const std::string &name, + const std::string &backend_name, + size_t num_workers, + size_t num_threads); + + void + destroy(); void fillRegList(nixl_xfer_dlist_t& dlist, nixlBasicDesc& desc) const; std::string getLocalMD() const; void loadRemoteMD(const std::string& remote_name); @@ -95,6 +109,7 @@ class TestErrorHandling : public testing::TestWithParam { bool dataCmp(const Agent& other) const; private: + std::string m_name; nixlBackendH* m_backend = nullptr; std::unique_ptr m_priv = nullptr; std::string m_MetaRemote; @@ -106,38 +121,58 @@ class TestErrorHandling : public testing::TestWithParam { BASIC_XFER, LOAD_REMOTE_THEN_FAIL, XFER_THEN_FAIL, + XFER_FAIL_RESTORE, + FAIL_AFTER_POST, }; TestErrorHandling(); template void testXfer(); private: - template bool isFailure(size_t iter); + template + bool + failBeforePost(size_t iter); + template + bool + failAfterPost(size_t iter); + template + bool + isFailure(size_t iter); template size_t numIter(); - void exchangeMetaData(); - nixlXferReqH* postXfer(enum nixl_xfer_op_t op, bool target_failure); + void + exchangeMetaData(); + template + std::variant + postXfer(enum nixl_xfer_op_t op, size_t iter); ScopedEnv m_env; Agent m_Initiator; Agent m_Target; std::string m_backend_name; + size_t numWorkers_; + size_t numThreads_; }; -void TestErrorHandling::Agent::init(const std::string& name, const std::string& backend_name) { +void +TestErrorHandling::Agent::init(const std::string &name, + const std::string &backend_name, + size_t num_workers, + size_t num_threads) { m_priv = std::make_unique(name, nixlAgentConfig(true)); // At the moment, only UCX backend is tested for error handling support. - m_backend = nixl::createUcxBackend(*m_priv, backend_name); + m_backend = nixl::createUcxBackend(*m_priv, backend_name, num_workers, num_threads); m_mem.init(m_backend); m_mem.fillData(); - EXPECT_EQ(NIXL_SUCCESS, - m_priv->registerMem(m_mem.m_dlist, &m_mem.m_params)); + EXPECT_EQ(NIXL_SUCCESS, m_priv->registerMem(m_mem.m_dlist, &m_mem.m_params)); } -void TestErrorHandling::Agent::destroy() { - m_MetaRemote.clear(); +void +TestErrorHandling::Agent::destroy() { m_priv->deregisterMem(m_mem.m_dlist, &m_mem.m_params); + m_priv->invalidateRemoteMD(m_MetaRemote); m_priv.reset(); + m_backend = nullptr; } void TestErrorHandling::Agent::fillRegList(nixl_xfer_dlist_t& dlist, @@ -152,7 +187,8 @@ std::string TestErrorHandling::Agent::getLocalMD() const { } void TestErrorHandling::Agent::loadRemoteMD(const std::string& remote_name) { - EXPECT_EQ(NIXL_SUCCESS, m_priv->loadRemoteMD(remote_name, m_MetaRemote)); + EXPECT_EQ(NIXL_SUCCESS, m_priv->loadRemoteMD(remote_name, m_MetaRemote)) + << "Agent " << m_name << " failed to load remote metadata"; } nixl_status_t @@ -168,19 +204,21 @@ TestErrorHandling::Agent::createXferReq(const nixl_xfer_op_t& op, } nixl_status_t -TestErrorHandling::Agent::postXferReq(nixlXferReqH* req_handle) const { +TestErrorHandling::Agent::postXferReq(nixlXferReqH *req_handle) const { return m_priv->postXferReq(req_handle); } nixl_status_t -TestErrorHandling::Agent::waitForCompletion(nixlXferReqH* req_handle) { +TestErrorHandling::Agent::waitForCompletion(nixlXferReqH *req_handle) { nixl_status_t status; do { status = m_priv->getXferStatus(req_handle); + EXPECT_NE(NIXL_ERR_NOT_POSTED, status); } while (status == NIXL_IN_PROG); m_priv->releaseXferReq(req_handle); + return status; } @@ -206,8 +244,10 @@ bool TestErrorHandling::Agent::dataCmp(const TestErrorHandling::Agent& other) co return m_mem.m_data == other.m_mem.m_data; } -TestErrorHandling::TestErrorHandling() : m_backend_name(GetParam()) -{ +TestErrorHandling::TestErrorHandling() + : m_backend_name(std::get<0>(GetParam())), + numWorkers_(std::get<1>(GetParam())), + numThreads_(std::get<2>(GetParam())) { m_env.addVar("UCX_RC_TIMEOUT", "100us"); m_env.addVar("UCX_RC_RETRY_COUNT", "4"); m_env.addVar("UCX_UD_TIMEOUT", "3s"); @@ -216,17 +256,36 @@ TestErrorHandling::TestErrorHandling() : m_backend_name(GetParam()) template void TestErrorHandling::testXfer() { - m_Initiator.init("initiator", m_backend_name); - m_Target.init("target", m_backend_name); + const std::string initiator_name = "initiator"; + const std::string target_name = "target"; + m_Initiator.init(initiator_name, m_backend_name, numWorkers_, numThreads_); + m_Target.init(target_name, m_backend_name, numWorkers_, numThreads_); exchangeMetaData(); for (size_t i = 0; i < numIter(); ++i) { - nixlXferReqH* req_handle = postXfer(op, isFailure(i)); - nixl_status_t status = m_Initiator.waitForCompletion(req_handle); + nixl_status_t status; + auto result = postXfer(op, i); + if (std::holds_alternative(result)) { + // Transfer completed immediately + status = std::get(result); + } else { + // Transfer was posted, wait for completion + nixlXferReqH *req_handle = std::get(result); + status = m_Initiator.waitForCompletion(req_handle); + } if (isFailure(i)) { - EXPECT_EQ(NIXL_ERR_REMOTE_DISCONNECT, status); + if (failBeforePost(i)) { + EXPECT_EQ(status, NIXL_ERR_REMOTE_DISCONNECT); + } else { + EXPECT_TRUE((status == NIXL_ERR_REMOTE_DISCONNECT) || (status == NIXL_SUCCESS)); + } + + if (test_type == TestType::XFER_FAIL_RESTORE) { + m_Target.init(target_name, m_backend_name, numWorkers_, numThreads_); + exchangeMetaData(); + } } else { EXPECT_EQ(NIXL_SUCCESS, status); EXPECT_EQ(NIXL_SUCCESS, m_Target.waitForNotif("notification")); @@ -240,27 +299,59 @@ void TestErrorHandling::testXfer() { switch (test_type) { case TestType::BASIC_XFER: + case TestType::XFER_FAIL_RESTORE: m_Target.destroy(); + m_Initiator.destroy(); + return; case TestType::LOAD_REMOTE_THEN_FAIL: case TestType::XFER_THEN_FAIL: + case TestType::FAIL_AFTER_POST: m_Initiator.destroy(); - break; - default: - EXPECT_TRUE(false) << "Invalid test type"; + return; } } template -bool TestErrorHandling::isFailure(size_t iter) { +bool +TestErrorHandling::failBeforePost(size_t iter) { switch (test_type) { - case TestType::BASIC_XFER: return false; - case TestType::LOAD_REMOTE_THEN_FAIL: return iter == 0; - case TestType::XFER_THEN_FAIL: return iter == 1; + case TestType::BASIC_XFER: + return false; + case TestType::LOAD_REMOTE_THEN_FAIL: + return iter == 0; + case TestType::XFER_THEN_FAIL: + case TestType::XFER_FAIL_RESTORE: + return iter == 1; + case TestType::FAIL_AFTER_POST: + return false; } } -template size_t TestErrorHandling::numIter() { - return (test_type == TestType::XFER_THEN_FAIL) ? 2 : 1; +template +bool +TestErrorHandling::failAfterPost(size_t iter) { + return (test_type == TestType::FAIL_AFTER_POST) && (iter == 1); +} + +template +bool +TestErrorHandling::isFailure(size_t iter) { + return failBeforePost(iter) || failAfterPost(iter); +} + +template +size_t +TestErrorHandling::numIter() { + switch (test_type) { + case TestType::BASIC_XFER: + case TestType::LOAD_REMOTE_THEN_FAIL: + return 1; + case TestType::XFER_THEN_FAIL: + case TestType::FAIL_AFTER_POST: + return 2; + case TestType::XFER_FAIL_RESTORE: + return 3; + } } void TestErrorHandling::exchangeMetaData() { @@ -268,8 +359,9 @@ void TestErrorHandling::exchangeMetaData() { m_Target.loadRemoteMD(m_Initiator.getLocalMD()); } -nixlXferReqH* -TestErrorHandling::postXfer(enum nixl_xfer_op_t op, bool target_failure) { +template +std::variant +TestErrorHandling::postXfer(enum nixl_xfer_op_t op, size_t iter) { EXPECT_TRUE(op == NIXL_WRITE || op == NIXL_READ); nixlBasicDesc sReq_src; @@ -281,31 +373,30 @@ TestErrorHandling::postXfer(enum nixl_xfer_op_t op, bool target_failure) { m_Target.fillRegList(rReq_descs, rReq_dst); nixlXferReqH* req_handle; - nixl_status_t status; - - status = m_Initiator.createXferReq(op, sReq_descs, rReq_descs, req_handle); + nixl_status_t status = m_Initiator.createXferReq(op, sReq_descs, rReq_descs, req_handle); EXPECT_EQ(NIXL_SUCCESS, status) - << "createXferReq failed with unexpected error: " - << nixlEnumStrings::statusStr(status); + << "createXferReq failed with unexpected error: " << nixlEnumStrings::statusStr(status); - if (target_failure) { + if (failBeforePost(iter)) { m_Target.destroy(); } status = m_Initiator.postXferReq(req_handle); - if (target_failure) { - // If the target is destroyed, the transfer may fail immediately - // or later - EXPECT_TRUE((status == NIXL_ERR_REMOTE_DISCONNECT) || - (status == NIXL_IN_PROG)); - } else { - EXPECT_LE(0, status) << "status: " - << nixlEnumStrings::statusStr(status); + + if (failAfterPost(iter)) { + m_Target.destroy(); } + if (isFailure(iter) && (status == NIXL_ERR_REMOTE_DISCONNECT)) { + // failed handle destroyed on post + return status; + } + + EXPECT_LE(0, status) << "status: " << nixlEnumStrings::statusStr(status); return req_handle; } + TEST_P(TestErrorHandling, BasicXfer) { testXfer(); testXfer(); @@ -321,6 +412,22 @@ TEST_P(TestErrorHandling, XferThenFail) { testXfer(); } -INSTANTIATE_TEST_SUITE_P(UCX, TestErrorHandling, testing::Values("UCX", "UCX_MO")); +TEST_P(TestErrorHandling, XferFailRestore) { + testXfer(); + testXfer(); +} + +TEST_P(TestErrorHandling, XferPostThenFail) { + testXfer(); + testXfer(); +} + +INSTANTIATE_TEST_SUITE_P(ucx, TestErrorHandling, testing::Values(std::make_tuple("UCX", 1, 0))); +INSTANTIATE_TEST_SUITE_P(ucx_mo, + TestErrorHandling, + testing::Values(std::make_tuple("UCX_MO", 1, 0))); +INSTANTIATE_TEST_SUITE_P(ucx_threadpool, + TestErrorHandling, + testing::Values(std::make_tuple("UCX", 2, 1))); } // namespace gtest diff --git a/test/gtest/main.cpp b/test/gtest/main.cpp index acfdfa6810..a04a56f507 100644 --- a/test/gtest/main.cpp +++ b/test/gtest/main.cpp @@ -15,7 +15,12 @@ * limitations under the License. */ #include "plugin_manager.h" +#include "common.h" #include +#include +#include +#include +#include namespace gtest { std::vector SplitWithDelimiter(const std::string &str, @@ -30,6 +35,19 @@ std::vector SplitWithDelimiter(const std::string &str, return tokens; } +void +ParseTcpPortRange(const std::string &arg) { + if (arg.find("--min-tcp-port=") == 0) { + const std::string min_port = SplitWithDelimiter(arg, '=').back(); + PortAllocator::instance().set_min_port(std::stoi(min_port)); + } + + if (arg.find("--max-tcp-port=") == 0) { + const std::string max_port = SplitWithDelimiter(arg, '=').back(); + PortAllocator::instance().set_max_port(std::stoi(max_port)); + } +} + void ParseArguments(int argc, char **argv) { for (int i = 1; i < argc; ++i) { if (std::string(argv[i]).find("--tests_plugin_dirs=") == 0) { @@ -42,6 +60,8 @@ void ParseArguments(int argc, char **argv) { } } } + + ParseTcpPortRange(argv[i]); } } diff --git a/test/gtest/meson.build b/test/gtest/meson.build index afa69af510..728c5ced6a 100644 --- a/test/gtest/meson.build +++ b/test/gtest/meson.build @@ -35,6 +35,10 @@ subdir('mocks') subdir('unit') subdir('plugins') +if ucx_gpu_device_api_available + subdir('device_api') +endif + plugin_dirs_arg = '--tests_plugin_dirs=' + mocks_dep.get_variable('path') cpp_flags = [] @@ -61,13 +65,24 @@ gtest_sources = [ 'test_transfer.cpp', 'metadata_exchange.cpp', 'common.cpp', - 'query_mem.cpp' + 'query_mem.cpp', + 'telemetry_test.cpp' ] + +if ucx_gpu_device_api_available + gtest_sources += device_api_test_sources + device_api_inc = [nixl_gpu_inc_dirs, include_directories('device_api')] + device_api_dep = ucx_dep +else + device_api_inc = [] + device_api_dep = [] +endif + test_exe = executable('gtest', sources : gtest_sources, - include_directories: [nixl_inc_dirs, utils_inc_dirs], + include_directories: [nixl_inc_dirs, utils_inc_dirs, device_api_inc], cpp_args : cpp_flags, - dependencies : [nixl_dep, cuda_dep, gtest_dep, gmock_dep, absl_strings_dep, absl_time_dep, file_utils_interface], + dependencies : [nixl_dep, nixl_common_dep, cuda_dep, device_api_dep, gtest_dep, gmock_dep, absl_strings_dep, absl_time_dep, file_utils_interface], link_with: [nixl_build_lib], install : true ) diff --git a/test/gtest/metadata_exchange.cpp b/test/gtest/metadata_exchange.cpp index 4699bf8044..9d86a85895 100644 --- a/test/gtest/metadata_exchange.cpp +++ b/test/gtest/metadata_exchange.cpp @@ -40,22 +40,6 @@ namespace gtest { namespace metadata_exchange { - -namespace { - -int getRandomPort() -{ - static constexpr int min_port = 10000; - static constexpr int max_port = 65535; - static std::random_device rd; - static std::mt19937 gen(rd()); - static std::uniform_int_distribution distr(min_port, max_port); - - return distr(gen); -} - -}; // unnamed namespace - class MemBuffer { public: MemBuffer(size_t size) : @@ -136,11 +120,9 @@ class MetadataExchangeTestFixture : public testing::Test { void SetUp() override { - int port_base = getRandomPort(); - // Create two agents for (int i = 0; i < AGENT_COUNT_; i++) { - int port = port_base + i; + const auto port = PortAllocator::next_tcp_port(); std::string name = "agent_" + std::to_string(i); nixlAgentConfig cfg(false, true, port, nixl_thread_sync_t::NIXL_THREAD_SYNC_STRICT); diff --git a/test/gtest/mocks/gmock_engine.cpp b/test/gtest/mocks/gmock_engine.cpp index 3dbd86dbc2..4dd01bff21 100644 --- a/test/gtest/mocks/gmock_engine.cpp +++ b/test/gtest/mocks/gmock_engine.cpp @@ -30,7 +30,6 @@ GMockBackendEngine::GMockBackendEngine() : nixlBackendEngine(&init_params) { ON_CALL(*this, supportsRemote()).WillByDefault(Return(true)); ON_CALL(*this, supportsLocal()).WillByDefault(Return(true)); ON_CALL(*this, supportsNotif()).WillByDefault(Return(true)); - ON_CALL(*this, supportsProgTh()).WillByDefault(Return(false)); ON_CALL(*this, getSupportedMems()).WillByDefault(Return(nixl_mem_list_t{DRAM_SEG})); ON_CALL(*this, registerMem(_, _, _)).WillByDefault(Return(NIXL_SUCCESS)); ON_CALL(*this, deregisterMem(_)).WillByDefault(Return(NIXL_SUCCESS)); @@ -51,7 +50,6 @@ GMockBackendEngine::GMockBackendEngine() : nixlBackendEngine(&init_params) { ON_CALL(*this, loadLocalMD(_, _)).WillByDefault(Return(NIXL_SUCCESS)); ON_CALL(*this, getNotifs(_)).WillByDefault(Return(NIXL_SUCCESS)); ON_CALL(*this, genNotif(_, _)).WillByDefault(Return(NIXL_SUCCESS)); - ON_CALL(*this, progress()).WillByDefault(Return(0)); } void diff --git a/test/gtest/mocks/gmock_engine.h b/test/gtest/mocks/gmock_engine.h index 88e9c96040..8739e87e97 100644 --- a/test/gtest/mocks/gmock_engine.h +++ b/test/gtest/mocks/gmock_engine.h @@ -56,6 +56,9 @@ class GMockBackendEngine : public nixlBackendEngine { public: GMockBackendEngine(); + GMockBackendEngine(const nixlBackendInitParams *init_params) : nixlBackendEngine(init_params) {} + + void SetToParams(nixl_b_params_t ¶ms) const; static GMockBackendEngine * @@ -64,7 +67,6 @@ class GMockBackendEngine : public nixlBackendEngine { MOCK_METHOD(bool, supportsRemote, (), (const, override)); MOCK_METHOD(bool, supportsLocal, (), (const, override)); MOCK_METHOD(bool, supportsNotif, (), (const, override)); - MOCK_METHOD(bool, supportsProgTh, (), (const, override)); MOCK_METHOD(nixl_mem_list_t, getSupportedMems, (), (const, override)); MOCK_METHOD(nixl_status_t, registerMem, @@ -119,7 +121,6 @@ class GMockBackendEngine : public nixlBackendEngine { genNotif, (const std::string &remote_agent, const std::string &msg), (const, override)); - MOCK_METHOD(int, progress, (), (override)); }; } // namespace mocks diff --git a/test/gtest/mocks/meson.build b/test/gtest/mocks/meson.build index ac6767fcf3..03cbc160db 100644 --- a/test/gtest/mocks/meson.build +++ b/test/gtest/mocks/meson.build @@ -17,7 +17,7 @@ gtest_inc_dirs = include_directories('..') mock_backend_sources = ['gmock_engine.cpp', 'mock_backend_plugin.cpp', 'mock_backend_engine.cpp'] mock_backend_plugin = shared_library('MOCK_BACKEND', mock_backend_sources, - dependencies: [nixl_infra, gmock_dep], + dependencies: [nixl_infra, nixl_common_dep, gmock_dep], include_directories: [nixl_inc_dirs, utils_inc_dirs, gtest_inc_dirs], link_with : [ucx_backend_lib], name_prefix: 'libplugin_', diff --git a/test/gtest/mocks/mock_backend_engine.cpp b/test/gtest/mocks/mock_backend_engine.cpp index 020f6afffd..452783cc91 100644 --- a/test/gtest/mocks/mock_backend_engine.cpp +++ b/test/gtest/mocks/mock_backend_engine.cpp @@ -124,9 +124,4 @@ MockBackendEngine::genNotif(const std::string &remote_agent, const std::string & return gmock_backend_engine->genNotif(remote_agent, msg); } -int -MockBackendEngine::progress() { - sharedState++; - return gmock_backend_engine->progress(); -} } // namespace mocks diff --git a/test/gtest/mocks/mock_backend_engine.h b/test/gtest/mocks/mock_backend_engine.h index e89a129c2d..61f64d9f09 100644 --- a/test/gtest/mocks/mock_backend_engine.h +++ b/test/gtest/mocks/mock_backend_engine.h @@ -44,10 +44,6 @@ class MockBackendEngine : public nixlBackendEngine { assert(sharedState > 0); return gmock_backend_engine->supportsNotif(); } - bool supportsProgTh() const override { - assert(sharedState > 0); - return gmock_backend_engine->supportsProgTh(); - } nixl_mem_list_t getSupportedMems() const override { assert(sharedState > 0); return gmock_backend_engine->getSupportedMems(); @@ -90,7 +86,6 @@ class MockBackendEngine : public nixlBackendEngine { nixl_status_t getNotifs(notif_list_t ¬if_list) override; nixl_status_t genNotif(const std::string &remote_agent, const std::string &msg) const override; - int progress() override; private: // This represents an engine shared state that is read in every const method and modified in non-cost ones diff --git a/test/gtest/plugins/meson.build b/test/gtest/plugins/meson.build index 3f1842a6a4..e9e11b4c4d 100644 --- a/test/gtest/plugins/meson.build +++ b/test/gtest/plugins/meson.build @@ -20,7 +20,7 @@ if not aws_s3.found() endif plugins_test_exe = executable('plugins_gtest', - sources : ['../main.cpp', 'obj_plugin.cpp'], + sources : ['../main.cpp', '../common.cpp', 'obj_plugin.cpp'], include_directories: [nixl_inc_dirs, utils_inc_dirs, plugins_inc_dirs, '.'], dependencies : [nixl_dep, gtest_dep, absl_strings_dep, absl_time_dep, plugin_deps, obj_backend_interface], diff --git a/test/gtest/plugins/transfer_handler.h b/test/gtest/plugins/transfer_handler.h index 9c13b2f9e5..d2d8c8074d 100644 --- a/test/gtest/plugins/transfer_handler.h +++ b/test/gtest/plugins/transfer_handler.h @@ -190,8 +190,6 @@ template class transferHandler { while (ret == NIXL_IN_PROG && absl::Now() < end_time) { ret = srcBackendEngine_->checkXfer(handle); ASSERT_TRUE(ret == NIXL_SUCCESS || ret == NIXL_IN_PROG); - - if (dstBackendEngine_->supportsProgTh()) dstBackendEngine_->progress(); } NIXL_INFO << "\nTransfer complete"; @@ -220,7 +218,6 @@ template class transferHandler { while (num_notifs == 0 && absl::Now() < end_time) { ASSERT_EQ(dstBackendEngine_->getNotifs(target_notifs), NIXL_SUCCESS); num_notifs = target_notifs.size(); - if (srcBackendEngine_->supportsProgTh()) srcBackendEngine_->progress(); } NIXL_INFO << "\nNotification transfer complete"; diff --git a/test/gtest/telemetry_test.cpp b/test/gtest/telemetry_test.cpp new file mode 100644 index 0000000000..cb2d0529be --- /dev/null +++ b/test/gtest/telemetry_test.cpp @@ -0,0 +1,527 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "telemetry.h" +#include "telemetry_event.h" +#include "nixl_types.h" +#include "common.h" +#include "backend/backend_engine.h" +#include "mocks/gmock_engine.h" + +namespace fs = std::filesystem; +constexpr char TELEMETRY_ENABLED_VAR[] = "NIXL_TELEMETRY_ENABLE"; +constexpr char TELEMETRY_DIR_VAR[] = "NIXL_TELEMETRY_DIR"; + +// Custom mock backend class for testing backend telemetry events +class telemetryTestBackend : public mocks::GMockBackendEngine { +public: + telemetryTestBackend(const nixlBackendInitParams *init_params) + : mocks::GMockBackendEngine(init_params) {} + + void + addTestTelemetryEvent(const std::string &event_name, uint64_t value) { + addTelemetryEvent(event_name, value); + } +}; + +class telemetryTest : public ::testing::Test { +protected: + void + SetUp() override { + testDir_ = "/tmp/telemetry_test_files"; + testFile_ = testDir_.string() + "/test_telemetry"; + try { + if (!fs::exists(testDir_)) { + fs::create_directory(testDir_); + } + } + catch (const fs::filesystem_error &e) { + throw std::runtime_error("Could not create the directory for telemetry test."); + } + + envHelper_.addVar(TELEMETRY_ENABLED_VAR, "y"); + envHelper_.addVar(TELEMETRY_DIR_VAR, testDir_.string()); + } + + void + TearDown() override { + envHelper_.popVar(); + envHelper_.popVar(); + if (fs::exists(testDir_)) { + try { + fs::remove_all(testDir_); + } + catch (const fs::filesystem_error &e) { + // ignore can fail due to nsf + } + } + } + + void + validateState() { + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + EXPECT_EQ(buffer->version(), TELEMETRY_VERSION); + EXPECT_EQ(buffer->capacity(), capacity_); + EXPECT_EQ(buffer->size(), size_); + EXPECT_EQ(buffer->empty(), size_ == 0); + EXPECT_EQ(buffer->full(), size_ == capacity_); + } + + fs::path testDir_; + std::string testFile_; + gtest::ScopedEnv envHelper_; + size_t capacity_ = 4096; + size_t size_ = 0; + size_t readPos_ = 0; + size_t writePos_ = 0; + size_t mask_ = 4096 - 1; + backend_map_t backendMap_; +}; + +TEST_F(telemetryTest, BasicInitialization) { + EXPECT_NO_THROW({ + nixlTelemetry telemetry(testFile_, backendMap_); + validateState(); + }); +} + +TEST_F(telemetryTest, InitializationWithEmptyFileName) { + EXPECT_THROW({ nixlTelemetry telemetry("", backendMap_); }, std::invalid_argument); +} + +TEST_F(telemetryTest, CustomBufferSize) { + auto tmp_capacity = capacity_; + capacity_ = 32; + envHelper_.addVar(TELEMETRY_BUFFER_SIZE_VAR, "32"); + + EXPECT_NO_THROW({ + nixlTelemetry telemetry(testFile_, backendMap_); + validateState(); + }); + capacity_ = tmp_capacity; + envHelper_.popVar(); +} + +TEST_F(telemetryTest, InvalidBufferSize) { + envHelper_.addVar(TELEMETRY_BUFFER_SIZE_VAR, "0"); + + EXPECT_THROW({ nixlTelemetry telemetry(testFile_, backendMap_); }, std::invalid_argument); + envHelper_.popVar(); + envHelper_.addVar(TELEMETRY_BUFFER_SIZE_VAR, "1023"); + EXPECT_THROW({ nixlTelemetry telemetry(testFile_, backendMap_); }, std::invalid_argument); + envHelper_.popVar(); +} + +// Test transfer bytes tracking +TEST_F(telemetryTest, TransferBytesTracking) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + nixlTelemetry telemetry(testFile_, backendMap_); + + EXPECT_NO_THROW(telemetry.updateTxBytes(1024)); + EXPECT_NO_THROW(telemetry.updateRxBytes(1024)); + EXPECT_NO_THROW(telemetry.updateTxRequestsNum(1)); + EXPECT_NO_THROW(telemetry.updateRxRequestsNum(1)); + EXPECT_NO_THROW(telemetry.updateErrorCount(nixl_status_t::NIXL_ERR_BACKEND)); + EXPECT_NO_THROW(telemetry.updateMemoryRegistered(1024)); + EXPECT_NO_THROW(telemetry.updateMemoryDeregistered(1024)); + EXPECT_NO_THROW(telemetry.addXferTime(std::chrono::microseconds(100), true, 2000)); + + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + EXPECT_EQ(buffer->size(), 10); + EXPECT_EQ(buffer->version(), TELEMETRY_VERSION); + EXPECT_EQ(buffer->capacity(), capacity_); + EXPECT_EQ(buffer->empty(), false); + EXPECT_EQ(buffer->full(), false); + nixlTelemetryEvent event; + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_bytes"); + EXPECT_EQ(event.value_, 1024); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_rx_bytes"); + EXPECT_EQ(event.value_, 1024); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_requests_num"); + EXPECT_EQ(event.value_, 1); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_rx_requests_num"); + EXPECT_EQ(event.value_, 1); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, + nixlEnumStrings::statusStr(nixl_status_t::NIXL_ERR_BACKEND).c_str()); + EXPECT_EQ(event.value_, 1); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_memory_registered"); + EXPECT_EQ(event.value_, 1024); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_memory_deregistered"); + EXPECT_EQ(event.value_, 1024); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_xfer_time"); + EXPECT_EQ(event.value_, 100); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_bytes"); + EXPECT_EQ(event.value_, 2000); + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_requests_num"); + EXPECT_EQ(event.value_, 1); + envHelper_.popVar(); +} + +TEST_F(telemetryTest, TelemetryEventStructure) { + nixlTelemetryEvent event1( + 1234567890, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER, "test_event", 42); + + EXPECT_EQ(event1.timestampUs_, 1234567890); + EXPECT_EQ(event1.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER); + EXPECT_EQ(event1.value_, 42); + EXPECT_STREQ(event1.eventName_, "test_event"); +} + +TEST_F(telemetryTest, ShortRunInterval) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + + EXPECT_NO_THROW({ nixlTelemetry telemetry(testFile_, backendMap_); }); + envHelper_.popVar(); +} + +TEST_F(telemetryTest, LargeRunInterval) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "10000"); + + EXPECT_NO_THROW({ nixlTelemetry telemetry(testFile_, backendMap_); }); + envHelper_.popVar(); +} + +TEST_F(telemetryTest, BufferOverflowHandling) { + envHelper_.addVar(TELEMETRY_BUFFER_SIZE_VAR, "4"); + + nixlTelemetry telemetry(testFile_, backendMap_); + + for (int i = 0; i < 10; ++i) { + EXPECT_NO_THROW(telemetry.updateTxBytes(i * 100)); + } + + envHelper_.popVar(); +} + +TEST_F(telemetryTest, CustomTelemetryDirectory) { + fs::path custom_dir = testDir_ / "custom_telemetry"; + fs::create_directory(custom_dir); + envHelper_.addVar(TELEMETRY_DIR_VAR, custom_dir.string()); + + EXPECT_NO_THROW({ + fs::path telemetry_file = custom_dir / "test_telemetry"; + nixlTelemetry telemetry(telemetry_file.string(), backendMap_); + + EXPECT_TRUE(fs::exists(telemetry_file)); + }); + envHelper_.popVar(); +} + +TEST_F(telemetryTest, TelemetryCategoryStringConversion) { + for (int i = 0; i < static_cast(nixl_telemetry_category_t::NIXL_TELEMETRY_CUSTOM) + 1; + ++i) { + auto category = static_cast(i); + std::string category_str = nixlEnumStrings::telemetryCategoryStr(category); + EXPECT_FALSE(category_str.empty()); + EXPECT_NE(category_str, "BAD_CATEGORY"); + } + + auto invalid_category = static_cast(999); + std::string invalid_str = nixlEnumStrings::telemetryCategoryStr(invalid_category); + EXPECT_EQ(invalid_str, "BAD_CATEGORY"); +} + +// Test concurrent access (basic thread safety) +TEST_F(telemetryTest, ConcurrentAccess) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + testFile_ = testDir_.string() + "/test_concurrent_access"; + nixlTelemetry telemetry(testFile_, backendMap_); + + const int num_threads = 4; + const int operations_per_thread = 100; + + std::vector threads; + + // Create threads that perform different telemetry operations + for (int i = 0; i < num_threads; ++i) { + threads.emplace_back([&telemetry, i]() { + for (int j = 0; j < operations_per_thread; ++j) { + switch (i % 4) { + case 0: + telemetry.updateTxBytes(j * 100); + break; + case 1: + telemetry.updateRxBytes(j * 50); + break; + case 2: + telemetry.updateTxRequestsNum(j); + break; + case 3: + telemetry.updateRxRequestsNum(j); + break; + } + } + }); + } + + // Wait for all threads to complete + for (auto &thread : threads) { + thread.join(); + } + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + size_ = operations_per_thread * num_threads; + readPos_ = 0; + writePos_ = size_; + validateState(); + envHelper_.popVar(); +} + +TEST_F(telemetryTest, BackendTelemetryEventsCollection) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + nixlBackendInitParams init_params; + init_params.enableTelemetry_ = true; + nixl_b_params_t custom_params; + init_params.customParams = &custom_params; + // Create mock backends and add them to the backend map + auto mock_backend1 = std::make_unique(&init_params); + auto mock_backend2 = std::make_unique(&init_params); + + backendMap_["CUSTOM"] = mock_backend1.get(); + backendMap_["GPUNETIO"] = mock_backend2.get(); + + nixlTelemetry telemetry(testFile_, backendMap_); + + // Add some telemetry events to the backends + mock_backend1->addTestTelemetryEvent("backend1_event1", 100); + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + mock_backend1->addTestTelemetryEvent("backend1_event2", 200); + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + mock_backend2->addTestTelemetryEvent("backend2_event1", 300); + + // Wait for the telemetry to be written + std::this_thread::sleep_for(std::chrono::milliseconds(3)); + + // Verify that backend events are collected and written to buffer + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + + EXPECT_EQ(buffer->size(), 3); // Should have 3 backend events + + nixlTelemetryEvent event; + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend1_event1"); + EXPECT_EQ(event.value_, 100); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend1_event2"); + EXPECT_EQ(event.value_, 200); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend2_event1"); + EXPECT_EQ(event.value_, 300); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + envHelper_.popVar(); +} + +TEST_F(telemetryTest, BackendTelemetryEventsEmptyBackendMap) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + + // Create telemetry with empty backend map + backend_map_t empty_backend_map; + nixlTelemetry telemetry(testFile_, empty_backend_map); + + // Add some agent events + telemetry.updateTxBytes(1024); + telemetry.updateRxBytes(2048); + + // Wait for the telemetry to be written + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + + // Verify that only agent events are written (no backend events) + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + + EXPECT_EQ(buffer->size(), 2); // Only agent events + + nixlTelemetryEvent event; + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_bytes"); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_rx_bytes"); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER); + + envHelper_.popVar(); +} + +TEST_F(telemetryTest, BackendTelemetryEventsMixedWithAgentEvents) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + + nixlBackendInitParams init_params; + nixl_b_params_t custom_params; + init_params.customParams = &custom_params; + init_params.enableTelemetry_ = true; + auto mock_backend = std::make_unique(&init_params); + backendMap_["CUSTOM"] = mock_backend.get(); + + nixlTelemetry telemetry(testFile_, backendMap_); + + // Add agent events + telemetry.updateTxBytes(1024); + telemetry.updateErrorCount(nixl_status_t::NIXL_ERR_BACKEND); + + // Add backend events + mock_backend->addTestTelemetryEvent("backend_event1", 100); + mock_backend->addTestTelemetryEvent("backend_event2", 200); + + // Wait for the telemetry to be written + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + + // Verify that both agent and backend events are written + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + + EXPECT_EQ(buffer->size(), 4); // 2 agent events + 2 backend events + + nixlTelemetryEvent event; + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "agent_tx_bytes"); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_TRANSFER); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, + nixlEnumStrings::statusStr(nixl_status_t::NIXL_ERR_BACKEND).c_str()); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_ERROR); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend_event1"); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend_event2"); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + envHelper_.popVar(); +} + +TEST_F(telemetryTest, BackendTelemetryEventsDisabledTelemetry) { + // Disable telemetry + unsetenv(TELEMETRY_ENABLED_VAR); + + nixlBackendInitParams init_params; + nixl_b_params_t custom_params; + init_params.customParams = &custom_params; + init_params.enableTelemetry_ = false; + auto mock_backend = std::make_unique(&init_params); + backendMap_["CUSTOM"] = mock_backend.get(); + + nixlTelemetry telemetry(testFile_, backendMap_); + // Add backend events (should be ignored) + mock_backend->addTestTelemetryEvent("backend_event_disabled", 100); + + // Wait a bit + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + + // Verify that no events are written + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + + EXPECT_EQ(buffer->size(), 0); + + // Restore environment variable + setenv(TELEMETRY_ENABLED_VAR, "y", 1); +} + +TEST_F(telemetryTest, BackendTelemetryEventsMultipleBackends) { + envHelper_.addVar(TELEMETRY_RUN_INTERVAL_VAR, "1"); + + // Create multiple mock backends + nixlBackendInitParams init_params; + nixl_b_params_t custom_params; + init_params.customParams = &custom_params; + init_params.enableTelemetry_ = true; + auto mock_backend1 = std::make_unique(&init_params); + auto mock_backend2 = std::make_unique(&init_params); + auto mock_backend3 = std::make_unique(&init_params); + + backendMap_["CUSTOM"] = mock_backend1.get(); + backendMap_["GPUNETIO"] = mock_backend2.get(); + backendMap_["GDS_MT"] = mock_backend3.get(); + + nixlTelemetry telemetry(testFile_, backendMap_); + + // Add events to each backend + mock_backend1->addTestTelemetryEvent("backend1_event", 100); + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + mock_backend2->addTestTelemetryEvent("backend2_event", 200); + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + mock_backend3->addTestTelemetryEvent("backend3_event", 300); + + // Wait for the telemetry to be written + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + + // Verify that events from all backends are collected + auto path = fs::path(testFile_); + auto buffer = std::make_unique>( + path.string(), false, TELEMETRY_VERSION); + + EXPECT_EQ(buffer->size(), 3); + + nixlTelemetryEvent event; + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend1_event"); + EXPECT_EQ(event.value_, 100); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend2_event"); + EXPECT_EQ(event.value_, 200); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + buffer->pop(event); + EXPECT_STREQ(event.eventName_, "backend3_event"); + EXPECT_EQ(event.value_, 300); + EXPECT_EQ(event.category_, nixl_telemetry_category_t::NIXL_TELEMETRY_BACKEND); + + envHelper_.popVar(); +} diff --git a/test/gtest/test_transfer.cpp b/test/gtest/test_transfer.cpp index 0e66e3d02d..73abbdd421 100644 --- a/test/gtest/test_transfer.cpp +++ b/test/gtest/test_transfer.cpp @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -35,6 +36,8 @@ #include #endif +constexpr auto min_chrono_time = std::chrono::steady_clock::time_point::min(); + namespace gtest { class MemBuffer : std::shared_ptr { @@ -93,18 +96,25 @@ class MemBuffer : std::shared_ptr { const size_t size; }; -class TestTransfer : public testing::TestWithParam { +class TestTransfer : + // Tuple fields are: backend_name, enable_progress_thread, num_workers, num_threads + public testing::TestWithParam> { protected: - static nixlAgentConfig getConfig(int listen_port) - { - return nixlAgentConfig(true, listen_port > 0, listen_port, - nixl_thread_sync_t::NIXL_THREAD_SYNC_RW, 0, - 100000); + nixlAgentConfig + getConfig(int listen_port, bool capture_telemetry) { + return nixlAgentConfig(isProgressThreadEnabled(), + listen_port > 0, + listen_port, + nixl_thread_sync_t::NIXL_THREAD_SYNC_RW, + 1, + 0, + 100000, + capture_telemetry); } - static int getPort(int i) - { - return 9000 + i; + uint16_t + getPort(int i) const { + return ports.at(i); } nixl_b_params_t getBackendParams() @@ -112,27 +122,39 @@ class TestTransfer : public testing::TestWithParam { nixl_b_params_t params; if (getBackendName() == "UCX" || getBackendName() == "UCX_MO") { - params["num_workers"] = "2"; + params["num_workers"] = std::to_string(getNumWorkers()); + params["num_threads"] = std::to_string(getNumThreads()); + params["split_batch_size"] = "32"; } return params; } + void + addAgent(unsigned int agent_num, bool capture_telemetry = false) { + ports.push_back(PortAllocator::next_tcp_port()); + agents.emplace_back(std::make_unique( + getAgentName(agent_num), getConfig(getPort(agent_num), capture_telemetry))); + nixlBackendH *backend_handle = nullptr; + nixl_status_t status = + agents.back()->createBackend(getBackendName(), getBackendParams(), backend_handle); + ASSERT_EQ(status, NIXL_SUCCESS); + EXPECT_NE(backend_handle, nullptr); + backend_handles.push_back(backend_handle); + } + void SetUp() override { #ifdef HAVE_CUDA m_cuda_device = (cudaSetDevice(0) == cudaSuccess); #endif + // Disabling Telemetry until the corresponding test + env.addVar("NIXL_TELEMETRY_ENABLE", "n"); + // Create two agents for (size_t i = 0; i < 2; i++) { - agents.emplace_back(std::make_unique(getAgentName(i), - getConfig(getPort(i)))); - nixlBackendH *backend_handle = nullptr; - nixl_status_t status = agents.back()->createBackend( - getBackendName(), getBackendParams(), backend_handle); - ASSERT_EQ(status, NIXL_SUCCESS); - EXPECT_NE(backend_handle, nullptr); + addAgent(i); } } @@ -143,11 +165,26 @@ class TestTransfer : public testing::TestWithParam { std::string getBackendName() const { - return GetParam(); + return std::get<0>(GetParam()); } - static nixl_opt_args_t extra_params_ip(int remote) - { + bool + isProgressThreadEnabled() const { + return std::get<1>(GetParam()); + } + + size_t + getNumWorkers() const { + return std::get<2>(GetParam()); + } + + size_t + getNumThreads() const { + return std::get<3>(GetParam()); + } + + nixl_opt_args_t + extra_params_ip(int remote) { nixl_opt_args_t extra_params; extra_params.ipAddr = "127.0.0.1"; @@ -196,10 +233,11 @@ class TestTransfer : public testing::TestWithParam { return result; } - void exchangeMDIP() - { - for (size_t i = 0; i < agents.size(); i++) { - for (size_t j = 0; j < agents.size(); j++) { + void + exchangeMDIP(size_t start, size_t end) { + // Exchange metadata for the agents in the specified range using their IP + for (size_t i = start; i <= end; i++) { + for (size_t j = start; j <= end; j++) { if (i == j) { continue; } @@ -212,15 +250,15 @@ class TestTransfer : public testing::TestWithParam { } } - void exchangeMD() - { - // Connect the existing agents and exchange metadata - for (size_t i = 0; i < agents.size(); i++) { + void + exchangeMD(size_t start, size_t end) { + // Exchange metadata for the agents in the specified range + for (size_t i = start; i <= end; i++) { nixl_blob_t md; nixl_status_t status = agents[i]->getLocalMD(md); ASSERT_EQ(status, NIXL_SUCCESS); - for (size_t j = 0; j < agents.size(); j++) { + for (size_t j = start; j <= end; j++) { if (i == j) continue; std::string remote_agent_name; @@ -231,11 +269,11 @@ class TestTransfer : public testing::TestWithParam { } } - void invalidateMD() - { - // Disconnect the agents and invalidate remote metadata - for (size_t i = 0; i < agents.size(); i++) { - for (size_t j = 0; j < agents.size(); j++) { + void + invalidateMD(size_t start, size_t end) { + // Invalidate each other's metadata for the agents in the specified range + for (size_t i = start; i <= end; i++) { + for (size_t j = start; j < end; j++) { if (i == j) continue; nixl_status_t status = agents[j]->invalidateRemoteMD( @@ -245,28 +283,6 @@ class TestTransfer : public testing::TestWithParam { } } - void waitForXfer(nixlAgent &from, const std::string &from_name, - nixlAgent &to, nixlXferReqH *xfer_req) - { - nixl_notifs_t notif_map; - bool xfer_done; - do { - // progress on "from" agent while waiting for notification - nixl_status_t status = from.getXferStatus(xfer_req); - EXPECT_TRUE((status == NIXL_SUCCESS) || (status == NIXL_IN_PROG)); - xfer_done = (status == NIXL_SUCCESS); - - // Get notifications and progress all agents to avoid deadlocks - status = to.getNotifs(notif_map); - ASSERT_EQ(status, NIXL_SUCCESS); - } while (notif_map.empty() || !xfer_done); - - // Expect the notification from the right agent - auto ¬if_list = notif_map[from_name]; - EXPECT_EQ(notif_list.size(), 1u); - EXPECT_EQ(notif_list.front(), NOTIF_MSG); - } - void createRegisteredMem(nixlAgent& agent, size_t size, size_t count, nixl_mem_t mem_type, @@ -287,10 +303,11 @@ class TestTransfer : public testing::TestWithParam { agent.deregisterMem(desc_list); } - void verifyNotifs(nixlAgent &agent, const std::string &from_name, size_t expected_count) - { - nixl_notifs_t notif_map; - + void + verifyNotifs(nixlAgent &agent, + const std::string &from_name, + size_t expected_count, + nixl_notifs_t notif_map = {}) { for (int i = 0; i < retry_count; i++) { nixl_status_t status = agent.getNotifs(notif_map); ASSERT_EQ(status, NIXL_SUCCESS); @@ -318,14 +335,20 @@ class TestTransfer : public testing::TestWithParam { size_t num_threads) { const size_t total_notifs = repeat * num_threads; - exchangeMD(); + exchangeMD(0, 1); std::vector threads; + nixl_notifs_t notif_map; for (size_t thread = 0; thread < num_threads; ++thread) { threads.emplace_back([&]() { for (size_t i = 0; i < repeat; ++i) { nixl_status_t status = from.genNotif(to_name, NOTIF_MSG); ASSERT_EQ(status, NIXL_SUCCESS); + + if (!isProgressThreadEnabled()) { + ASSERT_EQ(NIXL_SUCCESS, from.getNotifs(notif_map)); + ASSERT_EQ(NIXL_SUCCESS, to.getNotifs(notif_map)); + } } }); } @@ -334,21 +357,27 @@ class TestTransfer : public testing::TestWithParam { thread.join(); } - verifyNotifs(to, from_name, total_notifs); - - invalidateMD(); + verifyNotifs(to, from_name, total_notifs, std::move(notif_map)); + invalidateMD(0, 1); } - void doTransfer(nixlAgent &from, const std::string &from_name, - nixlAgent &to, const std::string &to_name, size_t size, - size_t count, size_t repeat, size_t num_threads, - nixl_mem_t src_mem_type, - std::vector src_buffers, - nixl_mem_t dst_mem_type, - std::vector dst_buffers) - { + void + doTransfer(nixlAgent &from, + const std::string &from_name, + nixlAgent &to, + const std::string &to_name, + size_t size, + size_t count, + size_t repeat, + size_t num_threads, + nixl_mem_t src_mem_type, + std::vector src_buffers, + nixl_mem_t dst_mem_type, + std::vector dst_buffers, + nixl_status_t expected_telem_status = NIXL_ERR_NO_TELEMETRY) { std::mutex logger_mutex; std::vector threads; + nixl_notifs_t notif_map; for (size_t thread = 0; thread < num_threads; ++thread) { threads.emplace_back([&, thread]() { nixl_opt_args_t extra_params; @@ -376,6 +405,9 @@ class TestTransfer : public testing::TestWithParam { if (status == NIXL_SUCCESS) { break; } + if (!isProgressThreadEnabled()) { + ASSERT_EQ(NIXL_SUCCESS, to.getNotifs(notif_map)); + } std::this_thread::sleep_for(retry_timeout); } EXPECT_TRUE(status == NIXL_SUCCESS); @@ -391,6 +423,16 @@ class TestTransfer : public testing::TestWithParam { << "(" << bandwidth << " GB/s)"; } + nixl_xfer_telem_t telemetry; + status = from.getXferTelemetry(xfer_req, telemetry); + EXPECT_EQ(status, expected_telem_status); + if (expected_telem_status == NIXL_SUCCESS) { + EXPECT_TRUE(telemetry.startTime > min_chrono_time); + EXPECT_TRUE(telemetry.postDuration > chrono_period_us_t(0)); + EXPECT_TRUE(telemetry.xferDuration > chrono_period_us_t(0)); + EXPECT_TRUE(telemetry.xferDuration >= telemetry.postDuration); + } + status = from.releaseXferReq(xfer_req); EXPECT_EQ(status, NIXL_SUCCESS); }); @@ -400,9 +442,7 @@ class TestTransfer : public testing::TestWithParam { thread.join(); } - verifyNotifs(to, from_name, repeat * num_threads); - - invalidateMD(); + verifyNotifs(to, from_name, repeat * num_threads, std::move(notif_map)); } nixlAgent &getAgent(size_t idx) @@ -416,14 +456,21 @@ class TestTransfer : public testing::TestWithParam { } bool m_cuda_device = false; + gtest::ScopedEnv env; + std::vector backend_handles; private: static constexpr uint64_t DEV_ID = 0; static const std::string NOTIF_MSG; - static constexpr int retry_count{1000}; + // TODO: with error handling enabled by default we get poor performance with UCX1.18. + // Before we upgrade to UCX1.19, we need to temporarily increase the retry count, + // in order to pass threadpool tests. + // TODO: revert this to 1000 once we upgrade to UCX1.19. + static constexpr int retry_count{10000}; static constexpr std::chrono::milliseconds retry_timeout{1}; std::vector> agents; + std::vector ports; }; const std::string TestTransfer::NOTIF_MSG = "notification"; @@ -445,7 +492,7 @@ TEST_P(TestTransfer, RandomSizes) createRegisteredMem(getAgent(0), size, count, mem_type, src_buffers); createRegisteredMem(getAgent(1), size, count, mem_type, dst_buffers); - exchangeMD(); + exchangeMD(0, 1); doTransfer(getAgent(0), getAgentName(0), getAgent(1), @@ -458,6 +505,7 @@ TEST_P(TestTransfer, RandomSizes) src_buffers, mem_type, dst_buffers); + invalidateMD(0, 1); deregisterMem(getAgent(0), src_buffers, mem_type); deregisterMem(getAgent(1), dst_buffers, mem_type); } @@ -473,12 +521,13 @@ TEST_P(TestTransfer, remoteMDFromSocket) createRegisteredMem(getAgent(0), size, count, mem_type, src_buffers); createRegisteredMem(getAgent(1), size, count, mem_type, dst_buffers); - exchangeMDIP(); + exchangeMDIP(0, 1); doTransfer(getAgent(0), getAgentName(0), getAgent(1), getAgentName(1), size, count, 1, 1, mem_type, src_buffers, mem_type, dst_buffers); + invalidateMD(0, 1); deregisterMem(getAgent(0), src_buffers, mem_type); deregisterMem(getAgent(1), dst_buffers, mem_type); } @@ -512,7 +561,180 @@ TEST_P(TestTransfer, ListenerCommSize) { deregisterMem(getAgent(1), buffers, DRAM_SEG); } -INSTANTIATE_TEST_SUITE_P(ucx, TestTransfer, testing::Values("UCX")); -INSTANTIATE_TEST_SUITE_P(ucx_mo, TestTransfer, testing::Values("UCX_MO")); +TEST_P(TestTransfer, GetXferTelemetryFile) { + env.addVar("NIXL_TELEMETRY_ENABLE", "y"); + env.addVar("NIXL_TELEMETRY_DIR", "/tmp/"); + + // Create fresh agents that read the current env var and add them to the fixture + addAgent(2); + addAgent(3); + + constexpr size_t size = 1024; + constexpr size_t count = 1; + std::vector src_buffers, dst_buffers; + createRegisteredMem(getAgent(2), size, count, DRAM_SEG, src_buffers); + createRegisteredMem(getAgent(3), size, count, DRAM_SEG, dst_buffers); + + exchangeMD(2, 3); + doTransfer(getAgent(2), + getAgentName(2), + getAgent(3), + getAgentName(3), + size, + count, + 1, + 1, + DRAM_SEG, + src_buffers, + DRAM_SEG, + dst_buffers, + NIXL_SUCCESS); + + invalidateMD(2, 3); + deregisterMem(getAgent(2), src_buffers, DRAM_SEG); + deregisterMem(getAgent(3), dst_buffers, DRAM_SEG); +} + +TEST_P(TestTransfer, GetXferTelemetryAPI) { + // Enable telemetry without file output + env.addVar("NIXL_TELEMETRY_ENABLE", "y"); + + // Create fresh agents that read the current env var and add them to the fixture + addAgent(2); + addAgent(3); + + constexpr size_t size = 1024; + constexpr size_t count = 1; + std::vector src_buffers, dst_buffers; + createRegisteredMem(getAgent(2), size, count, DRAM_SEG, src_buffers); + createRegisteredMem(getAgent(3), size, count, DRAM_SEG, dst_buffers); + + exchangeMD(2, 3); + doTransfer(getAgent(2), + getAgentName(2), + getAgent(3), + getAgentName(3), + size, + count, + 1, + 1, + DRAM_SEG, + src_buffers, + DRAM_SEG, + dst_buffers, + NIXL_SUCCESS); + + invalidateMD(2, 3); + deregisterMem(getAgent(2), src_buffers, DRAM_SEG); + deregisterMem(getAgent(3), dst_buffers, DRAM_SEG); +} + +TEST_P(TestTransfer, GetXferTelemetryAPICfg) { + // Disable telemetry from env var but through config, expecting a warning + env.addVar("NIXL_TELEMETRY_ENABLE", "n"); + + // Create fresh agents that read the current env var and add them to the fixture + // with capture_telemetry set + addAgent(2, true); + addAgent(3, true); + + constexpr size_t size = 1024; + constexpr size_t count = 1; + std::vector src_buffers, dst_buffers; + createRegisteredMem(getAgent(2), size, count, DRAM_SEG, src_buffers); + createRegisteredMem(getAgent(3), size, count, DRAM_SEG, dst_buffers); + + exchangeMD(2, 3); + doTransfer(getAgent(2), + getAgentName(2), + getAgent(3), + getAgentName(3), + size, + count, + 1, + 1, + DRAM_SEG, + src_buffers, + DRAM_SEG, + dst_buffers, + NIXL_SUCCESS); + + invalidateMD(2, 3); + deregisterMem(getAgent(2), src_buffers, DRAM_SEG); + deregisterMem(getAgent(3), dst_buffers, DRAM_SEG); +} + + +TEST_P(TestTransfer, GetXferTelemetryDisabled) { + env.addVar("NIXL_TELEMETRY_ENABLE", "n"); + + // Create fresh agents that read the current env var and add them to the fixture + addAgent(2); + addAgent(3); + + constexpr size_t size = 512; + constexpr size_t count = 1; + std::vector src_buffers, dst_buffers; + createRegisteredMem(getAgent(2), size, count, DRAM_SEG, src_buffers); + createRegisteredMem(getAgent(3), size, count, DRAM_SEG, dst_buffers); + + exchangeMD(2, 3); + doTransfer(getAgent(2), + getAgentName(2), + getAgent(3), + getAgentName(3), + size, + count, + 1, + 1, + DRAM_SEG, + src_buffers, + DRAM_SEG, + dst_buffers, + NIXL_ERR_NO_TELEMETRY); + + invalidateMD(2, 3); + deregisterMem(getAgent(2), src_buffers, DRAM_SEG); + deregisterMem(getAgent(3), dst_buffers, DRAM_SEG); +} + +TEST_P(TestTransfer, PrepGpuSignal) { +#ifndef HAVE_UCX_GPU_DEVICE_API + GTEST_SKIP() << "UCX GPU device API not available, skipping test"; +#else + size_t gpu_signal_size = 0; + nixl_opt_args_t extra_params = {.backends = {backend_handles[0]}}; + nixl_status_t size_status = getAgent(0).getGpuSignalSize(gpu_signal_size, &extra_params); + ASSERT_EQ(size_status, NIXL_SUCCESS) << "getGpuSignalSize failed"; + ASSERT_GT(gpu_signal_size, 0) << "GPU signal size is 0"; + + // Allocate a buffer on the GPU with the size of the signal + std::vector signal_buffer; + createRegisteredMem(getAgent(0), gpu_signal_size, 1, VRAM_SEG, signal_buffer); + + auto signal_desc_list = makeDescList(signal_buffer, VRAM_SEG); + + nixl_status_t status = getAgent(0).prepGpuSignal(signal_desc_list, &extra_params); + + EXPECT_EQ(status, NIXL_SUCCESS) + << "prepGpuSignal returned unexpected status: " << nixlEnumStrings::statusStr(status); + + deregisterMem(getAgent(0), signal_buffer, VRAM_SEG); +#endif +} + +INSTANTIATE_TEST_SUITE_P(ucx, TestTransfer, testing::Values(std::make_tuple("UCX", true, 2, 0))); +INSTANTIATE_TEST_SUITE_P(ucx_no_pt, + TestTransfer, + testing::Values(std::make_tuple("UCX", false, 2, 0))); +INSTANTIATE_TEST_SUITE_P(ucx_threadpool, + TestTransfer, + testing::Values(std::make_tuple("UCX", true, 6, 4))); +INSTANTIATE_TEST_SUITE_P(ucx_threadpool_no_pt, + TestTransfer, + testing::Values(std::make_tuple("UCX", false, 6, 4))); +INSTANTIATE_TEST_SUITE_P(ucx_mo, + TestTransfer, + testing::Values(std::make_tuple("UCX_MO", true, 2, 0))); } // namespace gtest diff --git a/test/gtest/unit/agent/meson.build b/test/gtest/unit/agent/meson.build index 4f12e19e84..86cd03a3f0 100644 --- a/test/gtest/unit/agent/meson.build +++ b/test/gtest/unit/agent/meson.build @@ -19,5 +19,5 @@ agent_unit_test_dep = declare_dependency( 'agent.cpp', ], include_directories: [nixl_inc_dirs, gtest_inc_dirs], - dependencies: [gmock_dep], -) \ No newline at end of file + dependencies: [gmock_dep, nixl_common_dep], +) diff --git a/test/gtest/unit/obj/obj.cpp b/test/gtest/unit/obj/obj.cpp index 728adb0748..3b76f08f2e 100644 --- a/test/gtest/unit/obj/obj.cpp +++ b/test/gtest/unit/obj/obj.cpp @@ -310,7 +310,6 @@ TEST_F(objTestFixture, EngineInitialization) { EXPECT_TRUE(objEngine_->supportsLocal()); EXPECT_FALSE(objEngine_->supportsRemote()); EXPECT_FALSE(objEngine_->supportsNotif()); - EXPECT_FALSE(objEngine_->supportsProgTh()); // Verify that the executor was properly set on the mock S3 client by the engine constructor EXPECT_TRUE(mockS3Client_->hasExecutor()); diff --git a/test/nixl/agent_example.cpp b/test/nixl/agent_example.cpp index 3c23e1228f..e30a8f71fb 100644 --- a/test/nixl/agent_example.cpp +++ b/test/nixl/agent_example.cpp @@ -81,9 +81,6 @@ void test_side_perf(nixlAgent* A1, nixlAgent* A2, nixlBackendH* backend, nixlBac } } - assert (src_list.verifySorted() == true); - assert (dst_list.verifySorted() == true); - assert (mem_list1.descCount() == n_mems); assert (mem_list2.descCount() == n_mems); @@ -386,13 +383,13 @@ nixl_status_t sideXferTest(nixlAgent* A1, nixlAgent* A2, nixlXferReqH* src_handl test_side_perf(A1, A2, src_backend, dst_backend); - int n_bufs = 4; //must be even + int n_bufs = 32; // must be even size_t len = 1024; void* src_bufs[n_bufs], *dst_bufs[n_bufs]; nixl_reg_dlist_t mem_list1(DRAM_SEG), mem_list2(DRAM_SEG); nixl_xfer_dlist_t src_list(DRAM_SEG), dst_list(DRAM_SEG); - nixlBlobDesc src_desc[4], dst_desc[4]; + nixlBlobDesc src_desc[n_bufs], dst_desc[n_bufs]; for(int i = 0; i staticPlugs; std::set plugins = { - "UCX", "GDS", "POSIX", "UCX_MO", "MOCK_BACKEND", "GPUNETIO", "OBJ", "GDS_MT"}; + "UCX", "GDS", "POSIX", "UCX_MO", "MOCK_BACKEND", "GPUNETIO", "OBJ", "GDS_MT", "LIBFABRIC"}; if (argc > 1 && (std::string(argv[1]) == "-h" || std::string(argv[1]) == "--help")) { print_usage(argv[0]); diff --git a/test/python/desc_perf.py b/test/python/desc_perf.py index a1e8309875..e543c0d4b4 100755 --- a/test/python/desc_perf.py +++ b/test/python/desc_perf.py @@ -2,23 +2,14 @@ # SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. import time import nixl._utils as nixl_utils from nixl._api import nixl_agent +from nixl.logging import get_logger + +logger = get_logger(__name__) if __name__ == "__main__": desc_count = 24 * 64 * 1024 @@ -29,13 +20,14 @@ start_time = time.perf_counter() - descs = agent.get_xfer_descs(addr_list, "DRAM", True) + descs = agent.get_xfer_descs(addr_list, "DRAM") end_time = time.perf_counter() assert descs.descCount() == desc_count - print( - "Time per desc add in us:", (1000000.0 * (end_time - start_time)) / desc_count + logger.info( + "Time per desc add in us: %f", + (1000000.0 * (end_time - start_time)) / desc_count, ) nixl_utils.free_passthru(addr) diff --git a/test/python/prep_xfer_perf.py b/test/python/prep_xfer_perf.py index c462e1d607..a5e1508ed5 100755 --- a/test/python/prep_xfer_perf.py +++ b/test/python/prep_xfer_perf.py @@ -21,6 +21,9 @@ import nixl._utils as nixl_utils from nixl._api import nixl_agent, nixl_agent_config +from nixl.logging import get_logger + +logger = get_logger(__name__) def init_agent(): @@ -34,26 +37,24 @@ def prep_handles(agent: nixl_agent, xfer_dlist, reg_dlist, indices): xfer_dlist_trim = reg_dlist.trim() elapsed = time.perf_counter() - start assert xfer_dlist_trim.descCount() == xfer_dlist.descCount() - print(f"Trim nixlRegDList:\t{elapsed:.4f} sec") + logger.info("Trim nixlRegDList:\t%.4f sec", elapsed) start = time.perf_counter() assert agent.register_memory(reg_dlist) is not None elapsed = time.perf_counter() - start - print(f"register_memory:\t{elapsed:.4f} sec") + logger.info("register_memory:\t%.4f sec", elapsed) start = time.perf_counter() - local_prep_handle = agent.prep_xfer_dlist( - "NIXL_INIT_AGENT", xfer_dlist, "DRAM", False - ) + local_prep_handle = agent.prep_xfer_dlist("NIXL_INIT_AGENT", xfer_dlist, "DRAM") elapsed = time.perf_counter() - start assert local_prep_handle - print(f"prep_xfer_dlist INIT:\t{elapsed:.4f} sec") + logger.info("prep_xfer_dlist INIT:\t%.4f sec", elapsed) start = time.perf_counter() - remote_prep_handle = agent.prep_xfer_dlist("agent", xfer_dlist, "DRAM", False) + remote_prep_handle = agent.prep_xfer_dlist("agent", xfer_dlist, "DRAM") elapsed = time.perf_counter() - start assert remote_prep_handle - print(f"prep_xfer_dlist SELF:\t{elapsed:.4f} sec") + logger.info("prep_xfer_dlist SELF:\t%.4f sec", elapsed) start = time.perf_counter() xfer_handle = agent.make_prepped_xfer( @@ -61,33 +62,33 @@ def prep_handles(agent: nixl_agent, xfer_dlist, reg_dlist, indices): ) elapsed = time.perf_counter() - start assert xfer_handle - print(f"make_prepped_xfer:\t{elapsed:.4f} sec") + logger.info("make_prepped_xfer:\t%.4f sec", elapsed) return local_prep_handle, remote_prep_handle, xfer_handle def perf_test_list(num_descs: int, addr_base: int, length: int): - print("-" * 40) - print("Starting list test...") - print("-" * 40) + logger.info("-" * 40) + logger.info("Starting list test...") + logger.info("-" * 40) agent = init_agent() descs_list = [(addr_base + i * length, length, 0) for i in range(num_descs)] indices = list(range(num_descs)) start = time.perf_counter() - xfer_dlist = agent.get_xfer_descs(descs_list, "DRAM", False) + xfer_dlist = agent.get_xfer_descs(descs_list, "DRAM") elapsed = time.perf_counter() - start assert xfer_dlist.descCount() == num_descs - print(f"get_xfer_descs:\t\t{elapsed:.4f} sec") + logger.info("get_xfer_descs:\t\t%.4f sec", elapsed) blob_descs_list = [ (addr_base + i * length, length, 0, b"") for i in range(num_descs) ] start = time.perf_counter() - reg_dlist = agent.get_reg_descs(blob_descs_list, "DRAM", False) + reg_dlist = agent.get_reg_descs(blob_descs_list, "DRAM") elapsed = time.perf_counter() - start assert reg_dlist.descCount() == num_descs - print(f"get_reg_descs:\t\t{elapsed:.4f} sec") + logger.info("get_reg_descs:\t\t%.4f sec", elapsed) local_prep_handle, remote_prep_handle, xfer_handle = prep_handles( agent, xfer_dlist, reg_dlist, indices @@ -97,9 +98,9 @@ def perf_test_list(num_descs: int, addr_base: int, length: int): def perf_test_array(num_descs: int, addr_base: int, length: int): - print("-" * 40) - print("Starting array test...") - print("-" * 40) + logger.info("-" * 40) + logger.info("Starting array test...") + logger.info("-" * 40) agent = init_agent() descs_np = np.zeros((num_descs, 3), dtype=np.uint64) indices = np.arange(num_descs) @@ -108,16 +109,16 @@ def perf_test_array(num_descs: int, addr_base: int, length: int): descs_np[:, 2] = 0 start = time.perf_counter() - xfer_dlist = agent.get_xfer_descs(descs_np, "DRAM", False) + xfer_dlist = agent.get_xfer_descs(descs_np, "DRAM") elapsed = time.perf_counter() - start assert xfer_dlist.descCount() == num_descs - print(f"get_xfer_descs:\t\t{elapsed:.4f} sec") + logger.info("get_xfer_descs:\t\t%.4f sec", elapsed) start = time.perf_counter() - reg_dlist = agent.get_reg_descs(descs_np, "DRAM", False) + reg_dlist = agent.get_reg_descs(descs_np, "DRAM") elapsed = time.perf_counter() - start assert reg_dlist.descCount() == num_descs - print(f"get_reg_descs:\t\t{elapsed:.4f} sec") + logger.info("get_reg_descs:\t\t%.4f sec", elapsed) local_prep_handle, remote_prep_handle, xfer_handle = prep_handles( agent, xfer_dlist, reg_dlist, indices @@ -138,8 +139,7 @@ def perf_test_array(num_descs: int, addr_base: int, length: int): ) args = parser.parse_args() - print("Using NIXL Plugins from:") - print(os.environ["NIXL_PLUGIN_DIR"]) + logger.info("Using NIXL Plugins from:\n%s", os.environ["NIXL_PLUGIN_DIR"]) # Example using nixl_agent_config agent = init_agent() @@ -147,7 +147,9 @@ def perf_test_array(num_descs: int, addr_base: int, length: int): num_descs = 2**8 length = 1024 addr_base = nixl_utils.malloc_passthru(num_descs * length) - print(f"Performance test: Creating nixlXferDList with {num_descs} descriptors") + logger.info( + "Performance test: Creating nixlXferDList with %d descriptors", num_descs + ) if args.mode == "list": perf_test_list(num_descs, addr_base, length) diff --git a/test/python/test_nixl_api.py b/test/python/test_nixl_api.py index e7a69c62d5..58af32ff33 100644 --- a/test/python/test_nixl_api.py +++ b/test/python/test_nixl_api.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os import uuid import pytest @@ -123,6 +124,7 @@ def test_metadata_pass(two_ucx_agents): passed_name = agent2.add_remote_agent(agent1.get_agent_metadata()) assert passed_name == agent1.name.encode() + utils.free_passthru(addr) @pytest.mark.timeout(5) @@ -200,3 +202,109 @@ def test_incorrect_plugin_env(monkeypatch): with pytest.raises(RuntimeError): nixl_agent("bad env agent") + + +def test_get_xfer_telemetry(): + os.environ["NIXL_TELEMETRY_ENABLE"] = "y" + + agent1 = nixl_agent(str(uuid.uuid4())) + agent2 = nixl_agent(str(uuid.uuid4())) + + mem_size = 128 + addr1 = utils.malloc_passthru(mem_size) + addr2 = utils.malloc_passthru(mem_size) + + try: + reg1 = agent1.get_reg_descs([(addr1, mem_size, 0, "")], mem_type="DRAM") + reg2 = agent2.get_reg_descs([(addr2, mem_size, 0, "")], mem_type="DRAM") + agent1.register_memory(reg1) + agent2.register_memory(reg2) + + agent1.add_remote_agent(agent2.get_agent_metadata()) + src = agent1.get_xfer_descs( + [(addr1, mem_size // 2, 0), (addr1 + mem_size // 2, mem_size // 2, 0)], + mem_type="DRAM", + ) + dst = agent1.get_xfer_descs( + [(addr2, mem_size // 2, 0), (addr2 + mem_size // 2, mem_size // 2, 0)], + mem_type="DRAM", + ) + + handle = agent1.initialize_xfer("WRITE", src, dst, agent2.name) + st = agent1.transfer(handle) + assert st in ("DONE", "PROC") + + while True: + st = agent1.check_xfer_state(handle) + assert st in ("DONE", "PROC") + if st == "DONE": + break + + telem = agent1.get_xfer_telemetry(handle) + assert telem.descCount == 2 + assert telem.totalBytes == mem_size + assert telem.startTime > 0 + assert telem.postDuration > 0 + assert telem.xferDuration > 0 + assert telem.xferDuration >= telem.postDuration + + agent1.release_xfer_handle(handle) + finally: + utils.free_passthru(addr1) + utils.free_passthru(addr2) + os.environ.pop("NIXL_TELEMETRY_ENABLE") + + +def test_get_xfer_telemetry_cfg(): + os.environ["NIXL_TELEMETRY_ENABLE"] = "m" # invalid value not to enable + os.environ["NIXL_TELEMETRY_DIR"] = "/tmp/dummy" # to be ignored + + agent1 = nixl_agent( + str(uuid.uuid4()), nixl_conf=nixl_agent_config(capture_telemetry=True) + ) + agent2 = nixl_agent(str(uuid.uuid4())) + + mem_size = 128 + addr1 = utils.malloc_passthru(mem_size) + addr2 = utils.malloc_passthru(mem_size) + + try: + reg1 = agent1.get_reg_descs([(addr1, mem_size, 0, "")], mem_type="DRAM") + reg2 = agent2.get_reg_descs([(addr2, mem_size, 0, "")], mem_type="DRAM") + agent1.register_memory(reg1) + agent2.register_memory(reg2) + + agent1.add_remote_agent(agent2.get_agent_metadata()) + src = agent1.get_xfer_descs( + [(addr1, mem_size // 2, 0), (addr1 + mem_size // 2, mem_size // 2, 0)], + mem_type="DRAM", + ) + dst = agent1.get_xfer_descs( + [(addr2, mem_size // 2, 0), (addr2 + mem_size // 2, mem_size // 2, 0)], + mem_type="DRAM", + ) + + handle = agent1.initialize_xfer("WRITE", src, dst, agent2.name) + st = agent1.transfer(handle) + assert st in ("DONE", "PROC") + + while True: + st = agent1.check_xfer_state(handle) + assert st in ("DONE", "PROC") + if st == "DONE": + break + + telem = agent1.get_xfer_telemetry(handle) + assert telem.descCount == 2 + assert telem.totalBytes == mem_size + assert telem.startTime > 0 + assert telem.postDuration > 0 + assert telem.xferDuration > 0 + assert telem.xferDuration >= telem.postDuration + + agent1.release_xfer_handle(handle) + finally: + utils.free_passthru(addr1) + utils.free_passthru(addr2) + os.environ.pop("NIXL_TELEMETRY_ENABLE") + os.environ.pop("NIXL_TELEMETRY_DIR") diff --git a/test/python/test_nixl_bindings.py b/test/python/test_nixl_bindings.py index b73aede291..6aed0cfd9b 100644 --- a/test/python/test_nixl_bindings.py +++ b/test/python/test_nixl_bindings.py @@ -19,13 +19,16 @@ import nixl._bindings as nixl import nixl._utils as nixl_utils +from nixl.logging import get_logger + +logger = get_logger(__name__) # These should automatically be run by pytest because of function names def test_list(): descs = [(1000, 105, 0), (2000, 30, 0), (1010, 20, 0)] - test_list = nixl.nixlXferDList(nixl.DRAM_SEG, descs, False) + test_list = nixl.nixlXferDList(nixl.DRAM_SEG, descs) assert test_list.descCount() == 3 @@ -33,7 +36,7 @@ def test_list(): pickled_list = pickle.dumps(test_list) - print(pickled_list) + logger.info("Pickled list: %s", pickled_list) unpickled_list = pickle.loads(pickled_list) @@ -41,7 +44,7 @@ def test_list(): assert test_list.getType() == nixl.DRAM_SEG - print(test_list.descCount()) + logger.info("Descriptor count: %s", test_list.descCount()) assert test_list.descCount() == 3 test_list.remDesc(1) @@ -57,6 +60,7 @@ def test_list(): def test_agent(): + os.environ["NIXL_TELEMETRY_ENABLE"] = "y" name1 = "Agent1" name2 = "Agent2" @@ -74,10 +78,10 @@ def test_agent(): nixl_utils.ba_buf(addr1, size) - reg_list1 = nixl.nixlRegDList(nixl.DRAM_SEG, False) + reg_list1 = nixl.nixlRegDList(nixl.DRAM_SEG) reg_list1.addDesc((addr1, size, 0, "dead")) - reg_list2 = nixl.nixlRegDList(nixl.DRAM_SEG, False) + reg_list2 = nixl.nixlRegDList(nixl.DRAM_SEG) reg_list2.addDesc((addr2, size, 0, "dead")) ret = agent1.registerMem(reg_list1, [ucx1]) @@ -89,10 +93,8 @@ def test_agent(): meta1 = agent1.getLocalMD() meta2 = agent2.getLocalMD() - print("Agent1 MD: ") - print(meta1) - print("Agent2 MD: ") - print(meta2) + logger.info("Agent1 MD: \n%s", meta1) + logger.info("Agent2 MD: \n%s", meta2) ret_name = agent1.loadRemoteMD(meta2) assert ret_name.decode(encoding="UTF-8") == name2 @@ -102,29 +104,29 @@ def test_agent(): offset = 8 req_size = 8 - src_list = nixl.nixlXferDList(nixl.DRAM_SEG, False) + src_list = nixl.nixlXferDList(nixl.DRAM_SEG) src_list.addDesc((addr1 + offset, req_size, 0)) - dst_list = nixl.nixlXferDList(nixl.DRAM_SEG, False) + dst_list = nixl.nixlXferDList(nixl.DRAM_SEG) dst_list.addDesc((addr2 + offset, req_size, 0)) - print("Transfer from " + str(addr1 + offset) + " to " + str(addr2 + offset)) + logger.info("Transfer from %s to %s", str(addr1 + offset), str(addr2 + offset)) noti_str = "n\0tification" - print(noti_str) + logger.info("Notification string: %s", noti_str) - print(src_list) - print(dst_list) + logger.info("Source list: %s", src_list) + logger.info("Destination list: %s", dst_list) handle = agent1.createXferReq(nixl.NIXL_WRITE, src_list, dst_list, name2, noti_str) assert handle != 0 - print(handle) + logger.info("Transfer handle: %s", handle) status = agent1.postXferReq(handle) assert status == nixl.NIXL_SUCCESS or status == nixl.NIXL_IN_PROG - print("Transfer posted") + logger.info("Transfer posted") notifMap = {} @@ -139,10 +141,19 @@ def test_agent(): nixl_utils.verify_transfer(addr1 + offset, addr2 + offset, req_size) assert len(notifMap[name1]) == 1 - print(notifMap[name1][0]) + logger.info("Received notification: %s", notifMap[name1][0]) assert notifMap[name1][0] == noti_str.encode() - print("Transfer verified") + logger.info("Transfer verified") + + # Verify transfer telemetry + telem = agent1.getXferTelemetry(handle) + assert telem.descCount == 1 + assert telem.totalBytes == req_size + assert telem.startTime > 0 + assert telem.postDuration > 0 + assert telem.xferDuration > 0 + assert telem.xferDuration >= telem.postDuration agent1.releaseXferReq(handle) @@ -185,7 +196,7 @@ def test_query_mem(): params, mems = agent.getPluginParams("POSIX") backend = agent.createBackend("POSIX", params) - descs = nixl.nixlRegDList(nixl.FILE_SEG, False) + descs = nixl.nixlRegDList(nixl.FILE_SEG) # Test 1: Query with empty descriptor list try: @@ -193,8 +204,9 @@ def test_query_mem(): assert len(resp) == 0 except Exception as e: # Some backends might not support queryMem, which is okay - print( - f"queryMem with empty list failed (expected for some backends): {e}" + logger.exception( + "queryMem with empty list failed (expected for some backends): %s", + e, ) # Test 2: Query with actual file descriptors @@ -228,16 +240,19 @@ def test_query_mem(): except Exception as e: # Some backends might not support queryMem, which is okay - print(f"queryMem failed (expected for some backends): {e}") + logger.exception( + "queryMem failed (expected for some backends): %s", + e, + ) except Exception as e: - print(f"Backend creation failed: {e}") + logger.exception("Backend creation failed: %s", e) # Try MOCK_DRAM as fallback try: params, mems = agent.getPluginParams("MOCK_DRAM") backend = agent.createBackend("MOCK_DRAM", params) - print("Using MOCK_DRAM backend") + logger.info("Using MOCK_DRAM backend") except Exception as e2: - print(f"MOCK_DRAM also failed: {e2}") + logger.exception("MOCK_DRAM also failed: %s", e2) return finally: diff --git a/test/unit/plugins/hf3fs/nixl_hf3fs_mt_test.cpp b/test/unit/plugins/hf3fs/nixl_hf3fs_mt_test.cpp index 2fcb61fd3b..ab14b41288 100644 --- a/test/unit/plugins/hf3fs/nixl_hf3fs_mt_test.cpp +++ b/test/unit/plugins/hf3fs/nixl_hf3fs_mt_test.cpp @@ -41,6 +41,7 @@ namespace { constexpr size_t default_transfer_size = 1024 * 1024; // 1MB constexpr int default_write_iterations = 1; constexpr int default_read_iterations = 1; + constexpr int default_iopool_size = 64; constexpr char test_phrase[] = "NIXL HF3FS Multi-Thread Test Pattern 2025"; constexpr char test_file_name[] = "mt_testfile"; constexpr mode_t std_file_permissions = 0744; @@ -307,10 +308,11 @@ int main(int argc, char *argv[]) { int write_iterations = default_write_iterations; int read_iterations = default_read_iterations; std::string test_dir = default_test_files_dir_path; + int iopool_size = default_iopool_size; // Parse command line arguments int opt; - while ((opt = getopt(argc, argv, "t:n:s:w:r:d:h")) != -1) { + while ((opt = getopt(argc, argv, "t:n:s:w:r:d:h:i:")) != -1) { switch (opt) { case 't': num_threads = std::stoi(optarg); @@ -330,6 +332,9 @@ int main(int argc, char *argv[]) { case 'd': test_dir = optarg; break; + case 'i': + iopool_size = std::stoi(optarg); + break; case 'h': default: std::cout << absl::StrFormat( @@ -356,6 +361,9 @@ int main(int argc, char *argv[]) { std::cout << absl::StrFormat(" -d: Test directory (default: %s)", default_test_files_dir_path) << std::endl; + std::cout << absl::StrFormat(" -i: IO Pool Size (default: %s)", + default_iopool_size) + << std::endl; return (opt == 'h') ? 0 : 1; } } @@ -370,11 +378,12 @@ int main(int argc, char *argv[]) { // Initialize NIXL nixlAgentConfig cfg(true, false, 0, nixl_thread_sync_t::NIXL_THREAD_SYNC_STRICT); - nixl_b_params_t params; nixlAgent agent("HF3FSMultiThreadTester", cfg); // Create HF3FS backend nixlBackendH* hf3fs = nullptr; + nixl_b_params_t params; + params["iopool_size"] = std::to_string(iopool_size); nixl_status_t ret = agent.createBackend("HF3FS", params, hf3fs); if (ret != NIXL_SUCCESS) { std::cerr << absl::StrFormat("Error creating HF3FS backend: %d", ret) << std::endl; diff --git a/test/unit/plugins/mooncake/mooncake_backend_test.cpp b/test/unit/plugins/mooncake/mooncake_backend_test.cpp index 3a3c526fcb..845e87ce17 100644 --- a/test/unit/plugins/mooncake/mooncake_backend_test.cpp +++ b/test/unit/plugins/mooncake/mooncake_backend_test.cpp @@ -28,10 +28,11 @@ int gpu_id = 0; -static void checkCudaError(cudaError_t result, const char *message) { +static void +checkCudaError(cudaError_t result, const char *message) { if (result != cudaSuccess) { - std::cerr << message << " (Error code: " << result << " - " - << cudaGetErrorString(result) << ")" << std::endl; + std::cerr << message << " (Error code: " << result << " - " << cudaGetErrorString(result) + << ")" << std::endl; exit(EXIT_FAILURE); } } @@ -45,7 +46,8 @@ class testHndlIterator { bool set; bool prepare; bool release; - nixlBackendReqH* handle; + nixlBackendReqH *handle; + public: testHndlIterator(bool _reuse) { reuse = _reuse; @@ -65,7 +67,8 @@ class testHndlIterator { assert(!set); } - bool needPrep() { + bool + needPrep() { if (reuse) { if (!prepare) { return false; @@ -74,18 +77,20 @@ class testHndlIterator { return true; } - bool needRelease() { + bool + needRelease() { return release; } - void isLast() { + void + isLast() { if (reuse) { release = true; } } - void setHandle(nixlBackendReqH *_handle) - { + void + setHandle(nixlBackendReqH *_handle) { assert(!set); handle = _handle; set = true; @@ -94,31 +99,32 @@ class testHndlIterator { } } - void unsetHandle() { + void + unsetHandle() { assert(set); set = false; } - nixlBackendReqH *&getHandle() { + nixlBackendReqH *& + getHandle() { assert(set); return handle; } }; - -nixlBackendEngine *createEngine(std::string name, bool p_thread) -{ - nixlBackendEngine *mooncake; +nixlBackendEngine * +createEngine(std::string name, bool p_thread) { + nixlBackendEngine *mooncake; nixlBackendInitParams init; - nixl_b_params_t custom_params; + nixl_b_params_t custom_params; init.enableProgTh = p_thread; - init.pthrDelay = 100; - init.localAgent = name; + init.pthrDelay = 100; + init.localAgent = name; init.customParams = &custom_params; - init.type = "Mooncake"; + init.type = "Mooncake"; - mooncake = (nixlBackendEngine*) new nixlMooncakeEngine (&init); + mooncake = (nixlBackendEngine *)new nixlMooncakeEngine(&init); assert(!mooncake->getInitErr()); if (mooncake->getInitErr()) { std::cout << "Failed to initialize worker1" << std::endl; @@ -128,14 +134,14 @@ nixlBackendEngine *createEngine(std::string name, bool p_thread) return mooncake; } -void releaseEngine(nixlBackendEngine *mooncake) -{ +void +releaseEngine(nixlBackendEngine *mooncake) { delete mooncake; } -std::string memType2Str(nixl_mem_t mem_type) -{ - switch(mem_type) { +std::string +memType2Str(nixl_mem_t mem_type) { + switch (mem_type) { case DRAM_SEG: return std::string("DRAM"); case VRAM_SEG: @@ -153,9 +159,8 @@ std::string memType2Str(nixl_mem_t mem_type) #ifdef HAVE_CUDA -static int cudaQueryAddr(void *address, bool &is_dev, - CUdevice &dev, CUcontext &ctx) -{ +static int +cudaQueryAddr(void *address, bool &is_dev, CUdevice &dev, CUcontext &ctx) { CUmemorytype mem_type = CU_MEMORYTYPE_HOST; uint32_t is_managed = 0; #define NUM_ATTRS 4 @@ -183,14 +188,14 @@ static int cudaQueryAddr(void *address, bool &is_dev, #endif -void allocateBuffer(nixl_mem_t mem_type, int dev_id, size_t len, void* &addr) -{ - switch(mem_type) { +void +allocateBuffer(nixl_mem_t mem_type, int dev_id, size_t len, void *&addr) { + switch (mem_type) { case DRAM_SEG: addr = calloc(1, len); break; #ifdef HAVE_CUDA - case VRAM_SEG:{ + case VRAM_SEG: { bool is_dev; CUdevice dev; CUcontext ctx; @@ -199,7 +204,7 @@ void allocateBuffer(nixl_mem_t mem_type, int dev_id, size_t len, void* &addr) checkCudaError(cudaMalloc(&addr, len), "Failed to allocate CUDA buffer 0"); cudaQueryAddr(addr, is_dev, dev, ctx); std::cout << "CUDA addr: " << std::hex << addr << " dev=" << std::dec << dev - << " ctx=" << std::hex << ctx << std::dec << std::endl; + << " ctx=" << std::hex << ctx << std::dec << std::endl; break; } #endif @@ -210,9 +215,9 @@ void allocateBuffer(nixl_mem_t mem_type, int dev_id, size_t len, void* &addr) assert(addr); } -void releaseBuffer(nixl_mem_t mem_type, int dev_id, void* &addr) -{ - switch(mem_type) { +void +releaseBuffer(nixl_mem_t mem_type, int dev_id, void *&addr) { + switch (mem_type) { case DRAM_SEG: free(addr); break; @@ -228,9 +233,9 @@ void releaseBuffer(nixl_mem_t mem_type, int dev_id, void* &addr) } } -void doMemset(nixl_mem_t mem_type, int dev_id, void *addr, char byte, size_t len) -{ - switch(mem_type) { +void +doMemset(nixl_mem_t mem_type, int dev_id, void *addr, char byte, size_t len) { + switch (mem_type) { case DRAM_SEG: memset(addr, byte, len); break; @@ -246,9 +251,9 @@ void doMemset(nixl_mem_t mem_type, int dev_id, void *addr, char byte, size_t len } } -void *getValidationPtr(nixl_mem_t mem_type, void *addr, size_t len) -{ - switch(mem_type) { +void * +getValidationPtr(nixl_mem_t mem_type, void *addr, size_t len) { + switch (mem_type) { case DRAM_SEG: return addr; break; @@ -265,9 +270,9 @@ void *getValidationPtr(nixl_mem_t mem_type, void *addr, size_t len) } } -void *releaseValidationPtr(nixl_mem_t mem_type, void *addr) -{ - switch(mem_type) { +void * +releaseValidationPtr(nixl_mem_t mem_type, void *addr) { + switch (mem_type) { case DRAM_SEG: break; #ifdef HAVE_CUDA @@ -282,16 +287,16 @@ void *releaseValidationPtr(nixl_mem_t mem_type, void *addr) return NULL; } -void allocateWrongGPUTest(nixlBackendEngine* mooncake, int dev_id) -{ +void +allocateWrongGPUTest(nixlBackendEngine *mooncake, int dev_id) { nixlBlobDesc desc; - nixlBackendMD* md; - void* buf; + nixlBackendMD *md; + void *buf; allocateBuffer(VRAM_SEG, dev_id, desc.len, buf); desc.devId = dev_id; - desc.addr = (uint64_t) buf; + desc.addr = (uint64_t)buf; int ret = mooncake->registerMem(desc, VRAM_SEG, md); @@ -300,15 +305,19 @@ void allocateWrongGPUTest(nixlBackendEngine* mooncake, int dev_id) releaseBuffer(VRAM_SEG, dev_id, buf); } -void allocateAndRegister(nixlBackendEngine *mooncake, int dev_id, nixl_mem_t mem_type, - void* &addr, size_t len, nixlBackendMD* &md) -{ +void +allocateAndRegister(nixlBackendEngine *mooncake, + int dev_id, + nixl_mem_t mem_type, + void *&addr, + size_t len, + nixlBackendMD *&md) { nixlBlobDesc desc; allocateBuffer(mem_type, dev_id, len, addr); - desc.addr = (uintptr_t) addr; - desc.len = len; + desc.addr = (uintptr_t)addr; + desc.len = len; desc.devId = dev_id; int ret = mooncake->registerMem(desc, mem_type, md); @@ -316,73 +325,85 @@ void allocateAndRegister(nixlBackendEngine *mooncake, int dev_id, nixl_mem_t mem assert(ret == NIXL_SUCCESS); } -void deallocateAndDeregister(nixlBackendEngine *mooncake, int dev_id, nixl_mem_t mem_type, - void* &addr, nixlBackendMD* &md) -{ +void +deallocateAndDeregister(nixlBackendEngine *mooncake, + int dev_id, + nixl_mem_t mem_type, + void *&addr, + nixlBackendMD *&md) { mooncake->deregisterMem(md); releaseBuffer(mem_type, dev_id, addr); } -void loadRemote(nixlBackendEngine *mooncake, int dev_id, std::string agent, - nixl_mem_t mem_type, void *addr, size_t len, - nixlBackendMD* &lmd, nixlBackendMD* &rmd) -{ +void +loadRemote(nixlBackendEngine *mooncake, + int dev_id, + std::string agent, + nixl_mem_t mem_type, + void *addr, + size_t len, + nixlBackendMD *&lmd, + nixlBackendMD *&rmd) { nixlBlobDesc info; - info.addr = (uintptr_t) addr; - info.len = len; - info.devId = dev_id; + info.addr = (uintptr_t)addr; + info.len = len; + info.devId = dev_id; mooncake->getPublicData(lmd, info.metaInfo); // Not applicable to Mooncake backend // assert(info.metaInfo.size() > 0); // We get the data from the cetnral location and populate the backend, and receive remote_meta - int ret = mooncake->loadRemoteMD (info, mem_type, agent, rmd); + int ret = mooncake->loadRemoteMD(info, mem_type, agent, rmd); assert(NIXL_SUCCESS == ret); } -void populateDescs(nixl_meta_dlist_t &descs, int dev_id, void *addr, int desc_cnt, size_t desc_size, nixlBackendMD* &md) -{ - for(int i = 0; i < desc_cnt; i++) { +void +populateDescs(nixl_meta_dlist_t &descs, + int dev_id, + void *addr, + int desc_cnt, + size_t desc_size, + nixlBackendMD *&md) { + for (int i = 0; i < desc_cnt; i++) { nixlMetaDesc req; - req.addr = (uintptr_t) (((char*) addr) + i * desc_size); //random offset - req.len = desc_size; - req.devId = dev_id; + req.addr = (uintptr_t)(((char *)addr) + i * desc_size); // random offset + req.len = desc_size; + req.devId = dev_id; req.metadataP = md; descs.addDesc(req); } } -static string op2string(nixl_xfer_op_t op, bool hasNotif) -{ - if(op == NIXL_READ && !hasNotif) - return string("READ"); - if(op == NIXL_WRITE && !hasNotif) - return string("WRITE"); - if(op == NIXL_READ && hasNotif) - return string("READ/NOTIF"); - if(op == NIXL_WRITE && hasNotif) - return string("WRITE/NOTIF"); +static string +op2string(nixl_xfer_op_t op, bool hasNotif) { + if (op == NIXL_READ && !hasNotif) return string("READ"); + if (op == NIXL_WRITE && !hasNotif) return string("WRITE"); + if (op == NIXL_READ && hasNotif) return string("READ/NOTIF"); + if (op == NIXL_WRITE && hasNotif) return string("WRITE/NOTIF"); return string("ERR-OP"); } - -void performTransfer(nixlBackendEngine *mooncake1, nixlBackendEngine *mooncake2, - nixl_meta_dlist_t &req_src_descs, - nixl_meta_dlist_t &req_dst_descs, - void* addr1, void* addr2, size_t len, - nixl_xfer_op_t op, - testHndlIterator &hiter, - bool progress, bool use_notif) -{ +void +performTransfer(nixlBackendEngine *mooncake1, + nixlBackendEngine *mooncake2, + nixl_meta_dlist_t &req_src_descs, + nixl_meta_dlist_t &req_dst_descs, + void *addr1, + void *addr2, + size_t len, + nixl_xfer_op_t op, + testHndlIterator &hiter, + bool progress, + bool use_notif) { int ret2; nixl_status_t ret3; void *chkptr1, *chkptr2; - std::string remote_agent ("Agent2"); + std::string remote_agent("Agent2"); - if(mooncake1 == mooncake2) remote_agent = "Agent1"; + if (mooncake1 == mooncake2) remote_agent = "Agent1"; std::string test_str("test"); std::cout << "\t" << op2string(op, use_notif) << " from " << addr1 << " to " << addr2 << "\n"; @@ -397,26 +418,25 @@ void performTransfer(nixlBackendEngine *mooncake1, nixlBackendEngine *mooncake2, // Also maybe we would remove the WRITE and let the backend class decide the op if (hiter.needPrep()) { nixlBackendReqH *new_handle = nullptr; - ret3 = mooncake1->prepXfer(op, req_src_descs, req_dst_descs, remote_agent, new_handle, &opt_args); + ret3 = mooncake1->prepXfer( + op, req_src_descs, req_dst_descs, remote_agent, new_handle, &opt_args); assert(ret3 == NIXL_SUCCESS); hiter.setHandle(new_handle); } nixlBackendReqH *&handle = hiter.getHandle(); ret3 = mooncake1->postXfer(op, req_src_descs, req_dst_descs, remote_agent, handle, &opt_args); - assert( ret3 == NIXL_SUCCESS || ret3 == NIXL_IN_PROG); + assert(ret3 == NIXL_SUCCESS || ret3 == NIXL_IN_PROG); if (ret3 == NIXL_SUCCESS) { - cout << "\t\tWARNING: Tansfer request completed immediately - no testing non-inline path" << endl; + cout << "\t\tWARNING: Tansfer request completed immediately - no testing non-inline path" + << endl; } else { cout << "\t\tNOTE: Testing non-inline Transfer path!" << endl; - while(ret3 == NIXL_IN_PROG) { + while (ret3 == NIXL_IN_PROG) { ret3 = mooncake1->checkXfer(handle); - if(progress){ - mooncake2->progress(); - } - assert( ret3 == NIXL_SUCCESS || ret3 == NIXL_IN_PROG); + assert(ret3 == NIXL_SUCCESS || ret3 == NIXL_IN_PROG); } } @@ -425,19 +445,16 @@ void performTransfer(nixlBackendEngine *mooncake1, nixlBackendEngine *mooncake2, mooncake1->releaseReqH(handle); } - if(use_notif) { - /* Test notification path */ + if (use_notif) { + /* Test notification path */ notif_list_t target_notifs; cout << "\t\tChecking notification flow: " << flush; ret2 = 0; - while(ret2 == 0){ + while (ret2 == 0) { ret3 = mooncake2->getNotifs(target_notifs); ret2 = target_notifs.size(); - if(progress){ - mooncake1->progress(); - } assert(ret3 == NIXL_SUCCESS); } @@ -455,8 +472,8 @@ void performTransfer(nixlBackendEngine *mooncake1, nixlBackendEngine *mooncake2, chkptr2 = getValidationPtr(req_dst_descs.getType(), addr2, len); // Perform correctness check. - for(size_t i = 0; i < len; i++){ - assert( ((uint8_t*) chkptr1)[i] == ((uint8_t*) chkptr2)[i]); + for (size_t i = 0; i < len; i++) { + assert(((uint8_t *)chkptr1)[i] == ((uint8_t *)chkptr2)[i]); } releaseValidationPtr(req_src_descs.getType(), chkptr1); @@ -465,14 +482,13 @@ void performTransfer(nixlBackendEngine *mooncake1, nixlBackendEngine *mooncake2, cout << "OK" << endl; } -void test_intra_agent_transfer(bool p_thread, nixlBackendEngine *mooncake, nixl_mem_t mem_type) -{ +void +test_intra_agent_transfer(bool p_thread, nixlBackendEngine *mooncake, nixl_mem_t mem_type) { std::cout << std::endl << std::endl; std::cout << "****************************************************" << std::endl; - std::cout << " Intra-agent memory transfer test: " - << "P-Thr=" << (p_thread ? "ON" : "OFF") << ", " << memType2Str(mem_type) - << std::endl; + std::cout << " Intra-agent memory transfer test: " << "P-Thr=" << (p_thread ? "ON" : "OFF") + << ", " << memType2Str(mem_type) << std::endl; std::cout << "****************************************************" << std::endl; std::cout << std::endl << std::endl; @@ -483,11 +499,11 @@ void test_intra_agent_transfer(bool p_thread, nixlBackendEngine *mooncake, nixl_ assert(mooncake->supportsLocal()); - //connection info is still a string + // connection info is still a string std::string conn_info1; ret1 = mooncake->getConnInfo(conn_info1); assert(ret1 == NIXL_SUCCESS); - ret1 = mooncake->loadRemoteConnInfo (agent1, conn_info1); + ret1 = mooncake->loadRemoteConnInfo(agent1, conn_info1); assert(ret1 == NIXL_SUCCESS); std::cout << "Local connection complete\n"; @@ -503,48 +519,63 @@ void test_intra_agent_transfer(bool p_thread, nixlBackendEngine *mooncake, nixl_ allocateAndRegister(mooncake, 0, mem_type, addr1, len, lmd1); allocateAndRegister(mooncake, 0, mem_type, addr2, len, lmd2); - //string descs unnecessary, convert meta locally - nixlBackendMD* rmd2; - ret1 = mooncake->loadLocalMD (lmd2, rmd2); + // string descs unnecessary, convert meta locally + nixlBackendMD *rmd2; + ret1 = mooncake->loadLocalMD(lmd2, rmd2); assert(ret1 == NIXL_SUCCESS); - nixl_meta_dlist_t req_src_descs (mem_type); + nixl_meta_dlist_t req_src_descs(mem_type); populateDescs(req_src_descs, 0, addr1, desc_cnt, desc_size, lmd1); - nixl_meta_dlist_t req_dst_descs (mem_type); + nixl_meta_dlist_t req_dst_descs(mem_type); populateDescs(req_dst_descs, 0, addr2, desc_cnt, desc_size, rmd2); - nixl_xfer_op_t ops[] = { NIXL_READ, NIXL_WRITE }; - bool use_notifs[] = { false }; // Mooncake transfer engine doesn't support notifs + nixl_xfer_op_t ops[] = {NIXL_READ, NIXL_WRITE}; + bool use_notifs[] = {true, false}; - for (size_t i = 0; i < sizeof(ops)/sizeof(ops[i]); i++) { + for (size_t i = 0; i < sizeof(ops) / sizeof(ops[i]); i++) { - for(bool use_notif : use_notifs) { - cout << endl << op2string(ops[i], use_notif) << " test (" << iter << ") iterations" <unloadMD (rmd2); + mooncake->unloadMD(rmd2); deallocateAndDeregister(mooncake, 0, mem_type, addr1, lmd1); deallocateAndDeregister(mooncake, 0, mem_type, addr2, lmd2); mooncake->disconnect(agent1); } -void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, - nixlBackendEngine *mooncake1, nixl_mem_t src_mem_type, int src_dev_id, - nixlBackendEngine *mooncake2, nixl_mem_t dst_mem_type, int dst_dev_id) -{ +void +test_inter_agent_transfer(bool p_thread, + bool reuse_hndl, + nixlBackendEngine *mooncake1, + nixl_mem_t src_mem_type, + int src_dev_id, + nixlBackendEngine *mooncake2, + nixl_mem_t dst_mem_type, + int dst_dev_id) { int ret; int iter = 10; @@ -553,8 +584,8 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, std::cout << " Inter-agent memory transfer test " << std::endl; std::cout << " P-Thr=" << (p_thread ? "ON" : "OFF") << std::endl; std::cout << " Handler-reuse=" << (reuse_hndl ? "ON" : "OFF") << std::endl; - std::cout << " (" << memType2Str(src_mem_type) << " -> " - << memType2Str(dst_mem_type) << ")" << std::endl; + std::cout << " (" << memType2Str(src_mem_type) << " -> " << memType2Str(dst_mem_type) + << ")" << std::endl; std::cout << "****************************************************" << std::endl; std::cout << std::endl << std::endl; @@ -572,7 +603,7 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, assert(ret == NIXL_SUCCESS); // We assumed we put them to central location and now receiving it on the other process - ret = mooncake1->loadRemoteConnInfo (agent2, conn_info2); + ret = mooncake1->loadRemoteConnInfo(agent2, conn_info2); assert(ret == NIXL_SUCCESS); // TODO: Causes race condition - investigate conn management implementation @@ -580,6 +611,18 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, std::cout << "Synchronous handshake complete\n"; + std::string test_str("test"); + mooncake1->genNotif(agent2, test_str); + int ret_gen = 0; + notif_list_t target_notif_gen; + while (ret_gen == 0) { + int ret3_gen = mooncake2->getNotifs(target_notif_gen); + ret_gen = target_notif_gen.size(); + assert(ret3_gen == NIXL_SUCCESS); + } + assert(target_notif_gen.front().second == test_str); + cout << "\t\tGenNotify Data verification success!" << flush; + // Number of transfer descriptors int desc_cnt = 64; // Size of a single descriptor @@ -592,43 +635,53 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, allocateAndRegister(mooncake2, dst_dev_id, dst_mem_type, addr2, len, lmd2); nixlBackendMD *rmd1 /*, *rmd2*/; - loadRemote(mooncake1, dst_dev_id, agent2, dst_mem_type, addr2, len, lmd2, rmd1); - //loadRemote(mooncake2, src_dev_id, agent1, src_mem_type, addr1, len, lmd1, rmd2); + loadRemote(mooncake1, dst_dev_id, agent2, dst_mem_type, addr2, len, lmd2, rmd1); + // loadRemote(mooncake2, src_dev_id, agent1, src_mem_type, addr1, len, lmd1, rmd2); - nixl_meta_dlist_t req_src_descs (src_mem_type); + nixl_meta_dlist_t req_src_descs(src_mem_type); populateDescs(req_src_descs, src_dev_id, addr1, desc_cnt, desc_size, lmd1); - nixl_meta_dlist_t req_dst_descs (dst_mem_type); + nixl_meta_dlist_t req_dst_descs(dst_mem_type); populateDescs(req_dst_descs, dst_dev_id, addr2, desc_cnt, desc_size, rmd1); - nixl_xfer_op_t ops[] = { NIXL_READ, NIXL_WRITE }; - bool use_notifs[] = { false }; // Mooncake transfer engine doesn't support notifs + nixl_xfer_op_t ops[] = {NIXL_READ, NIXL_WRITE}; + bool use_notifs[] = {true, false}; - for (size_t i = 0; i < sizeof(ops)/sizeof(ops[i]); i++) { + for (size_t i = 0; i < sizeof(ops) / sizeof(ops[i]); i++) { - for(bool use_notif : use_notifs) { - cout << endl << op2string(ops[i], use_notif) << " test (" << iter << ") iterations" <unloadMD (rmd1); - //mooncake2->unloadMD (rmd2); + mooncake1->unloadMD(rmd1); + // mooncake2->unloadMD (rmd2); // Release memory regions deallocateAndDeregister(mooncake1, src_dev_id, src_mem_type, addr1, lmd1); @@ -638,17 +691,17 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, mooncake1->disconnect(agent2); // TODO: Causes race condition - investigate conn management implementation - //mooncake2->disconnect(agent1); + // mooncake2->disconnect(agent1); } -int main() -{ +int +main() { bool thread_on[2] = {false, true}; - nixlBackendEngine *mooncake[2][2] = { 0 }; + nixlBackendEngine *mooncake[2][2] = {0}; // Allocate Mooncake engines - for(int i = 0; i < 2; i++) { - for(int j = 0; j < 2; j++) { + for (int i = 0; i < 2; i++) { + for (int j = 0; j < 2; j++) { std::stringstream s; s << "Agent" << (j + 1); mooncake[i][j] = createEngine(s.str(), thread_on[i]); @@ -656,7 +709,7 @@ int main() } #ifdef HAVE_CUDA - int dev_ids[2] = { 0 , 0 }; + int dev_ids[2] = {0, 0}; int n_vram_dev; if (cudaGetDeviceCount(&n_vram_dev) != cudaSuccess) { std::cout << "Call to cudaGetDeviceCount failed, assuming 0 devices"; @@ -669,10 +722,13 @@ int main() dev_ids[0] = 0; } #endif + test_inter_agent_transfer( + thread_on[0], false, mooncake[0][0], DRAM_SEG, 0, mooncake[0][1], DRAM_SEG, 0); - for(int i = 0; i < 2; i++) { - //Test local memory to local memory transfer - test_intra_agent_transfer(thread_on[i], mooncake[i][0], DRAM_SEG); + for (int i = 0; i < 2; i++) { + // Test local memory to local memory transfer + // std::cout << "thread_on" < 0) { test_intra_agent_transfer(thread_on[i], mooncake[i][0], VRAM_SEG); @@ -680,43 +736,61 @@ int main() #endif } - for(int i = 0; i < 2; i++) { - test_inter_agent_transfer(thread_on[i], false, - mooncake[i][0], DRAM_SEG, 0, - mooncake[i][1], DRAM_SEG, 0); - test_inter_agent_transfer(thread_on[i], true, - mooncake[i][0], DRAM_SEG, 0, - mooncake[i][1], DRAM_SEG, 0); + for (int i = 0; i < 2; i++) { + test_inter_agent_transfer( + thread_on[i], false, mooncake[i][0], DRAM_SEG, 0, mooncake[i][1], DRAM_SEG, 0); + test_inter_agent_transfer( + thread_on[i], true, mooncake[i][0], DRAM_SEG, 0, mooncake[i][1], DRAM_SEG, 0); #ifdef HAVE_CUDA if (n_vram_dev > 1) { - test_inter_agent_transfer(thread_on[i], false, - mooncake[i][0], VRAM_SEG, dev_ids[0], - mooncake[i][1], VRAM_SEG, dev_ids[1]); - test_inter_agent_transfer(thread_on[i], true, - mooncake[i][0], VRAM_SEG, dev_ids[0], - mooncake[i][1], VRAM_SEG, dev_ids[1]); - test_inter_agent_transfer(thread_on[i], true, - mooncake[i][0], DRAM_SEG, dev_ids[0], - mooncake[i][1], VRAM_SEG, dev_ids[1]); - test_inter_agent_transfer(thread_on[i], true, - mooncake[i][0], VRAM_SEG, dev_ids[0], - mooncake[i][1], DRAM_SEG, dev_ids[1]); + test_inter_agent_transfer(thread_on[i], + false, + mooncake[i][0], + VRAM_SEG, + dev_ids[0], + mooncake[i][1], + VRAM_SEG, + dev_ids[1]); + test_inter_agent_transfer(thread_on[i], + true, + mooncake[i][0], + VRAM_SEG, + dev_ids[0], + mooncake[i][1], + VRAM_SEG, + dev_ids[1]); + test_inter_agent_transfer(thread_on[i], + true, + mooncake[i][0], + DRAM_SEG, + dev_ids[0], + mooncake[i][1], + VRAM_SEG, + dev_ids[1]); + test_inter_agent_transfer(thread_on[i], + true, + mooncake[i][0], + VRAM_SEG, + dev_ids[0], + mooncake[i][1], + DRAM_SEG, + dev_ids[1]); } #endif } #ifdef HAVE_CUDA if (n_vram_dev > 1) { - //Test if registering on a different GPU fails correctly - allocateWrongGPUTest(mooncake[0][0], 1); - std::cout << "Verified registration on wrong GPU fails correctly\n"; - } + // Test if registering on a different GPU fails correctly + allocateWrongGPUTest(mooncake[0][0], 1); + std::cout << "Verified registration on wrong GPU fails correctly\n"; + } #endif // Deallocate Mooncake engines - for(int i = 0; i < 2; i++) { - for(int j = 0; j < 2; j++) { + for (int i = 0; i < 2; i++) { + for (int j = 0; j < 2; j++) { releaseEngine(mooncake[i][j]); } } diff --git a/test/unit/plugins/ucx/ucx_backend_multi.cpp b/test/unit/plugins/ucx/ucx_backend_multi.cpp index bc538514e7..cf3cdab0f7 100644 --- a/test/unit/plugins/ucx/ucx_backend_multi.cpp +++ b/test/unit/plugins/ucx/ucx_backend_multi.cpp @@ -32,7 +32,6 @@ void test_thread(int id) { nixlBackendInitParams init_params; nixl_b_params_t custom_params; - nixlBackendEngine* ucx; nixl_status_t ret; std::string my_name("Agent1"); @@ -50,9 +49,9 @@ void test_thread(int id) std::cout << my_name << " Started\n"; - ucx = (nixlBackendEngine*) new nixlUcxEngine (&init_params); + auto ucx = nixlUcxEngine::create(init_params).release(); - if(!USE_PTHREAD) ucx->progress(); + if (!USE_PTHREAD) ucx->progress(); ucx->getConnInfo(conn_info[id]); @@ -71,7 +70,7 @@ void test_thread(int id) done[id] = true; while(!done[!id]) - if(!USE_PTHREAD && id) ucx->progress(); + if (!USE_PTHREAD && id) ucx->progress(); std::cout << "Thread passed with id " << id << "\n"; @@ -83,7 +82,7 @@ void test_thread(int id) //wait for other while(!disconnect[!id]); - if(!USE_PTHREAD) ucx->progress(); + if (!USE_PTHREAD) ucx->progress(); std::cout << "Thread disconnected with id " << id << "\n"; diff --git a/test/unit/plugins/ucx/ucx_backend_test.cpp b/test/unit/plugins/ucx/ucx_backend_test.cpp index 28aa9c56cd..6f893b8009 100644 --- a/test/unit/plugins/ucx/ucx_backend_test.cpp +++ b/test/unit/plugins/ucx/ucx_backend_test.cpp @@ -106,10 +106,8 @@ class testHndlIterator { } }; - -nixlBackendEngine *createEngine(std::string name, bool p_thread) -{ - nixlBackendEngine *ucx; +nixlUcxEngine * +createEngine(std::string name, bool p_thread) { nixlBackendInitParams init; nixl_b_params_t custom_params; @@ -119,7 +117,7 @@ nixlBackendEngine *createEngine(std::string name, bool p_thread) init.customParams = &custom_params; init.type = "UCX"; - ucx = (nixlBackendEngine*) new nixlUcxEngine (&init); + auto ucx = nixlUcxEngine::create(init).release(); assert(!ucx->getInitErr()); if (ucx->getInitErr()) { std::cout << "Failed to initialize worker1" << std::endl; @@ -129,8 +127,8 @@ nixlBackendEngine *createEngine(std::string name, bool p_thread) return ucx; } -void releaseEngine(nixlBackendEngine *ucx) -{ +void +releaseEngine(nixlUcxEngine *ucx) { delete ucx; } @@ -283,8 +281,8 @@ void *releaseValidationPtr(nixl_mem_t mem_type, void *addr) return NULL; } -void allocateWrongGPUTest(nixlBackendEngine* ucx, int dev_id) -{ +void +allocateWrongGPUTest(nixlUcxEngine *ucx, int dev_id) { nixlBlobDesc desc; nixlBackendMD* md; void* buf; @@ -301,9 +299,13 @@ void allocateWrongGPUTest(nixlBackendEngine* ucx, int dev_id) releaseBuffer(VRAM_SEG, dev_id, buf); } -void allocateAndRegister(nixlBackendEngine *ucx, int dev_id, nixl_mem_t mem_type, - void* &addr, size_t len, nixlBackendMD* &md) -{ +void +allocateAndRegister(nixlUcxEngine *ucx, + int dev_id, + nixl_mem_t mem_type, + void *&addr, + size_t len, + nixlBackendMD *&md) { nixlBlobDesc desc; allocateBuffer(mem_type, dev_id, len, addr); @@ -317,17 +319,25 @@ void allocateAndRegister(nixlBackendEngine *ucx, int dev_id, nixl_mem_t mem_type assert(ret == NIXL_SUCCESS); } -void deallocateAndDeregister(nixlBackendEngine *ucx, int dev_id, nixl_mem_t mem_type, - void* &addr, nixlBackendMD* &md) -{ +void +deallocateAndDeregister(nixlUcxEngine *ucx, + int dev_id, + nixl_mem_t mem_type, + void *&addr, + nixlBackendMD *&md) { ucx->deregisterMem(md); releaseBuffer(mem_type, dev_id, addr); } -void loadRemote(nixlBackendEngine *ucx, int dev_id, std::string agent, - nixl_mem_t mem_type, void *addr, size_t len, - nixlBackendMD* &lmd, nixlBackendMD* &rmd) -{ +void +loadRemote(nixlUcxEngine *ucx, + int dev_id, + std::string agent, + nixl_mem_t mem_type, + void *addr, + size_t len, + nixlBackendMD *&lmd, + nixlBackendMD *&rmd) { nixlBlobDesc info; info.addr = (uintptr_t) addr; info.len = len; @@ -367,15 +377,18 @@ static string op2string(nixl_xfer_op_t op, bool hasNotif) return string("ERR-OP"); } - -void performTransfer(nixlBackendEngine *ucx1, nixlBackendEngine *ucx2, - nixl_meta_dlist_t &req_src_descs, - nixl_meta_dlist_t &req_dst_descs, - void* addr1, void* addr2, size_t len, - nixl_xfer_op_t op, - testHndlIterator &hiter, - bool progress, bool use_notif) -{ +void +performTransfer(nixlUcxEngine *ucx1, + nixlUcxEngine *ucx2, + nixl_meta_dlist_t &req_src_descs, + nixl_meta_dlist_t &req_dst_descs, + void *addr1, + void *addr2, + size_t len, + nixl_xfer_op_t op, + testHndlIterator &hiter, + bool progress, + bool use_notif) { int ret2; nixl_status_t ret3; void *chkptr1, *chkptr2; @@ -465,8 +478,8 @@ void performTransfer(nixlBackendEngine *ucx1, nixlBackendEngine *ucx2, cout << "OK" << endl; } -void test_intra_agent_transfer(bool p_thread, nixlBackendEngine *ucx, nixl_mem_t mem_type) -{ +void +test_intra_agent_transfer(bool p_thread, nixlUcxEngine *ucx, nixl_mem_t mem_type) { std::cout << std::endl << std::endl; std::cout << "****************************************************" << std::endl; @@ -541,10 +554,15 @@ void test_intra_agent_transfer(bool p_thread, nixlBackendEngine *ucx, nixl_mem_t ucx->disconnect(agent1); } -void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, - nixlBackendEngine *ucx1, nixl_mem_t src_mem_type, int src_dev_id, - nixlBackendEngine *ucx2, nixl_mem_t dst_mem_type, int dst_dev_id) -{ +void +test_inter_agent_transfer(bool p_thread, + bool reuse_hndl, + nixlUcxEngine *ucx1, + nixl_mem_t src_mem_type, + int src_dev_id, + nixlUcxEngine *ucx2, + nixl_mem_t dst_mem_type, + int dst_dev_id) { int ret; int iter = 10; @@ -674,7 +692,7 @@ void test_inter_agent_transfer(bool p_thread, bool reuse_hndl, int main() { bool thread_on[2] = {false, true}; - nixlBackendEngine *ucx[2][2] = { 0 }; + nixlUcxEngine *ucx[2][2] = {0}; // Allocate UCX engines for(int i = 0; i < 2; i++) { diff --git a/test/unit/plugins/ucx_mo/ucx_mo_backend_test.cpp b/test/unit/plugins/ucx_mo/ucx_mo_backend_test.cpp index 6ccb800c26..fd23ba7ec3 100644 --- a/test/unit/plugins/ucx_mo/ucx_mo_backend_test.cpp +++ b/test/unit/plugins/ucx_mo/ucx_mo_backend_test.cpp @@ -368,7 +368,7 @@ void performTransfer(nixlBackendEngine *ucx1, nixlBackendEngine *ucx2, while(status == NIXL_IN_PROG) { status = ucx1->checkXfer(handle); if(progress){ - ucx2->progress(); + ((nixlUcxMoEngine *)ucx2)->progress(); } assert( (NIXL_SUCCESS == status) || (NIXL_IN_PROG == status) ); } @@ -385,7 +385,7 @@ void performTransfer(nixlBackendEngine *ucx1, nixlBackendEngine *ucx2, status = ucx2->getNotifs(target_notifs); assert(NIXL_SUCCESS == status); if(progress){ - ucx1->progress(); + ((nixlUcxMoEngine *)ucx1)->progress(); } } @@ -526,7 +526,7 @@ void test_agent_transfer(bool p_thread, assert(NIXL_SUCCESS == status); if (!p_thread) { /* progress UCX1 as well */ - ucx1->progress(); + ((nixlUcxMoEngine *)ucx1)->progress(); } } diff --git a/test/unit/utils/libfabric/libfabric_topology_test.cpp b/test/unit/utils/libfabric/libfabric_topology_test.cpp new file mode 100644 index 0000000000..1352cf588e --- /dev/null +++ b/test/unit/utils/libfabric/libfabric_topology_test.cpp @@ -0,0 +1,66 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "libfabric/libfabric_topology.h" +#include "libfabric/libfabric_common.h" +#include "common/nixl_log.h" + +#ifdef CUDA_FOUND +#include +#endif + +int +main() { + NIXL_INFO << "=== Testing Libfabric Topology Implementation ==="; + try { + // Create topology instance - discovery happens automatically in constructor + NIXL_INFO << "1. Testing topology discovery..."; + nixlLibfabricTopology topology; + + NIXL_INFO << " SUCCESS: Topology discovery completed successfully"; + + // Print topology information + NIXL_INFO << "2. Topology Information:"; + topology.printTopologyInfo(); + + // Test GPU-specific queries only if GPUs are detected + int num_gpus = topology.getNumGpus(); + if (num_gpus > 0) { + NIXL_INFO << "3. Testing GPU-specific queries (detected " << num_gpus << " GPUs)..."; + int test_gpus = std::min(num_gpus, 3); // Test up to 3 GPUs or all available + for (int gpu_id = 0; gpu_id < test_gpus; ++gpu_id) { + auto gpu_devices = topology.getEfaDevicesForGpu(gpu_id); + std::string device_list; + for (const auto &device : gpu_devices) { + if (!device_list.empty()) device_list += " "; + device_list += device; + } + NIXL_INFO << " GPU " << gpu_id << " mapped to " << gpu_devices.size() + << " EFA devices: " << device_list; + } + } else { + NIXL_INFO << "3. Skipping GPU-specific tests (no GPUs detected)"; + } + } + catch (const std::exception &e) { + NIXL_ERROR << " Topology discovery failed: " << e.what(); + return 1; + } + NIXL_INFO << "=== Test completed successfully! ==="; + return 0; +} diff --git a/test/unit/utils/libfabric/meson.build b/test/unit/utils/libfabric/meson.build new file mode 100644 index 0000000000..85f8528b05 --- /dev/null +++ b/test/unit/utils/libfabric/meson.build @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2025 Amazon.com, Inc. and affiliates. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +libfabric_utils_dep = [ libfabric_dep, nixl_common_deps ] +libfabric_test_cpp_args = [] + +if cuda_dep.found() + libfabric_utils_dep += [ cuda_dep ] + libfabric_test_cpp_args += [ '-DCUDA_FOUND' ] +endif + +if get_option('buildtype') != 'release' + + libfabric_topology_test_bin = executable('libfabric_topology_test', + 'libfabric_topology_test.cpp', + dependencies: libfabric_utils_dep, + include_directories: [nixl_inc_dirs, utils_inc_dirs], + link_with: libfabric_utils_lib, + cpp_args: libfabric_test_cpp_args, + install: true) + +endif diff --git a/test/unit/utils/meson.build b/test/unit/utils/meson.build index a41a75d3e9..97cee34980 100644 --- a/test/unit/utils/meson.build +++ b/test/unit/utils/meson.build @@ -14,6 +14,9 @@ # limitations under the License. subdir('common') +if libfabric_dep.found() + subdir('libfabric') +endif subdir('serdes') subdir('stream') subdir('ucx')