Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
// SPDX-License-Identifier: Apache-2.0

import { InferenceParameters } from '@nemo/common/src/components/ModelSelectV2/InferenceParameters';
import type { InferenceParams } from '@nemo/sdk/generated/platform/schema';
import { fireEvent, render, screen } from '@testing-library/react';
import { useState } from 'react';

const renderComponent = (props: Partial<React.ComponentProps<typeof InferenceParameters>> = {}) => {
const onChange = vi.fn();
Expand Down Expand Up @@ -63,6 +65,21 @@ describe('InferenceParameters', () => {
);
});

it('lets the user clear the temperature field and type a new value', () => {
const StatefulHarness = () => {
const [value, setValue] = useState<Partial<InferenceParams>>({ temperature: 1 });
return <InferenceParameters value={value} onChange={setValue} />;
};
render(<StatefulHarness />);
const temperature = screen.getAllByRole('spinbutton')[0];

fireEvent.change(temperature, { target: { value: '' } });
expect(temperature).toHaveValue(null);

fireEvent.change(temperature, { target: { value: '0.7' } });
expect(temperature).toHaveValue(0.7);
});

it('disables all inputs when disabled', () => {
renderComponent({ disabled: true });
const inputs = screen.getAllByRole('spinbutton');
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -249,18 +249,43 @@ describe('SliderWithTextInput', () => {
/>
);

const slider = screen.getByRole('spinbutton');
const slider = screen.getByRole('slider');

// Test value above max
fireEvent.change(slider, { target: { value: '150' } });
fireEvent.input(slider, { target: { value: '150' } });
expect(mockOnChange).toHaveBeenCalledWith(90);

// Test value below min
fireEvent.change(slider, { target: { value: '5' } });
fireEvent.input(slider, { target: { value: '5' } });
expect(mockOnChange).toHaveBeenCalledWith(10);
});

it('should clamp values to min/max bounds when text input changes', () => {
it('should not clamp text input values while the user is still typing', () => {
const mockOnChange = vi.fn();
const field = createMockField(50, mockOnChange);
render(
<SliderWithTextInput
field={field}
defaultValue={25}
min={10}
max={90}
step={1}
disabled={false}
/>
);

const textInput = screen.getByRole('spinbutton');

// '5' is below min but is a valid prefix of '55', so it must survive.
fireEvent.change(textInput, { target: { value: '5' } });
expect(mockOnChange).not.toHaveBeenCalled();
expect(textInput).toHaveValue(5);

fireEvent.change(textInput, { target: { value: '55' } });
expect(mockOnChange).toHaveBeenCalledWith(55);
});

it('should clamp out-of-range text input values on blur', () => {
const mockOnChange = vi.fn();
const field = createMockField(50, mockOnChange);
render(
Expand All @@ -278,12 +303,52 @@ describe('SliderWithTextInput', () => {

// Test value above max
fireEvent.change(textInput, { target: { value: '150' } });
fireEvent.blur(textInput);
expect(mockOnChange).toHaveBeenCalledWith(90);

// Test value below min
fireEvent.change(textInput, { target: { value: '5' } });
fireEvent.blur(textInput);
expect(mockOnChange).toHaveBeenCalledWith(10);
});

it('should allow typing a value whose prefix is below min', () => {
const mockOnChange = vi.fn();
const field = createMockField(1, mockOnChange);
const { rerender } = render(
<SliderWithTextInput
field={field}
defaultValue={1}
min={0.1}
max={2}
step={0.1}
disabled={false}
/>
);

const textInput = screen.getByRole('spinbutton');

// Typing '0' must not be rewritten to the 0.1 minimum.
fireEvent.change(textInput, { target: { value: '0' } });
expect(mockOnChange).not.toHaveBeenCalled();
expect(textInput).toHaveValue(0);

fireEvent.change(textInput, { target: { value: '0.7' } });
expect(mockOnChange).toHaveBeenCalledWith(0.7);

rerender(
<SliderWithTextInput
field={createMockField(0.7, mockOnChange)}
defaultValue={1}
min={0.1}
max={2}
step={0.1}
disabled={false}
/>
);
fireEvent.blur(textInput);
expect(textInput).toHaveValue(0.7);
});
});

describe('Reset Functionality', () => {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
// SPDX-License-Identifier: Apache-2.0

import { toScientificNotation } from '@nemo/common/src/utils/formatters';
import { clamp } from '@nemo/common/src/utils/math';
import {
Block,
Button,
Expand All @@ -16,7 +17,7 @@ import {
Tooltip,
} from '@nvidia/foundations-react-core';
import { Info, RotateCcw } from 'lucide-react';
import { ComponentProps, ReactNode } from 'react';
import { ComponentProps, ReactNode, useState } from 'react';
Comment thread
steramae-nvidia marked this conversation as resolved.
import { FieldValues } from 'react-hook-form';

type SliderProps = ComponentProps<typeof Slider>;
Expand Down Expand Up @@ -57,31 +58,49 @@ export const SliderWithTextInput = ({
slotEnd,
displayName,
}: SliderWithTextInputProps) => {
const [draftValue, setDraftValue] = useState<string | null>(null);

const handleSliderChange = (newValue: number) => {
const clampedValue = Math.min(Math.max(newValue, min), max);
const clampedValue = clamp(newValue, min, max);
setDraftValue(null);
field.onChange(clampedValue);
attributes?.Slider?.onValueChange?.(clampedValue);
};
const handleTextInputChange = (newValue: string, event: React.ChangeEvent<HTMLInputElement>) => {
setDraftValue(newValue);
if (event.target.validity?.badInput) return;
if (newValue === '') {
field.onChange(undefined);
attributes?.TextInput?.onValueChange?.('', event);
return;
}
const numberValue = parseFloat(newValue);
if (Number.isNaN(numberValue)) return;
const clampedValue = Math.min(Math.max(numberValue, min), max);
if (numberValue < min || numberValue > max) return;
field.onChange(numberValue);
attributes?.TextInput?.onValueChange?.(numberValue.toString(), event);
};
const handleTextInputBlur = (event: React.FocusEvent<HTMLInputElement>) => {
const rawValue = draftValue;
setDraftValue(null);
attributes?.TextInput?.onBlur?.(event);
if (rawValue === null || rawValue === '') return;
const numberValue = parseFloat(rawValue);
if (Number.isNaN(numberValue)) return;
const clampedValue = clamp(numberValue, min, max);
if (clampedValue === field.value) return;
field.onChange(clampedValue);
attributes?.TextInput?.onValueChange?.(clampedValue.toString(), event);
attributes?.Slider?.onValueChange?.(clampedValue);
};
Comment thread
steramae-nvidia marked this conversation as resolved.
const handleReset = () => {
setDraftValue(null);
field.onChange(defaultValue);
attributes?.Slider?.onValueChange?.(defaultValue);
};
const fallback = defaultValue ?? min;
const isFieldValueNumber = typeof field.value === 'number' && !Number.isNaN(field.value);
const safeFieldValue = isFieldValueNumber ? field.value : fallback;
const textInputValue = isFieldValueNumber ? field.value.toString() : '';
const textInputValue = draftValue ?? (isFieldValueNumber ? field.value.toString() : '');

const stepMarkerClassNames =
'pb-5 [&_.nv-slider-step:first-of-type]:items-start [&_.nv-slider-step:last-of-type]:items-end';
Expand Down Expand Up @@ -156,6 +175,7 @@ export const SliderWithTextInput = ({
className={`${textInputWidth} h-[40px] shrink-0`}
{...attributes?.TextInput}
onValueChange={handleTextInputChange}
onBlur={handleTextInputBlur}
attributes={{
Input: {
'aria-label': `${id || 'slider'}_text_input`,
Expand Down
5 changes: 5 additions & 0 deletions web/packages/common/src/utils/math.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
// SPDX-License-Identifier: Apache-2.0

export const clamp = (value: number, min: number, max: number) =>
Math.min(Math.max(value, min), max);
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '3');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeLessThanOrEqual(2);
Expand Down Expand Up @@ -302,6 +303,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '-0.5');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeGreaterThanOrEqual(0);
Expand All @@ -323,6 +325,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '1.5');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeLessThanOrEqual(1);
Expand Down Expand Up @@ -428,6 +431,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '0.5');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeGreaterThanOrEqual(1);
Expand All @@ -452,6 +456,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '3');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeLessThanOrEqual(2);
Expand Down Expand Up @@ -742,6 +747,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '0');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeGreaterThanOrEqual(1);
Expand Down Expand Up @@ -775,6 +781,7 @@ describe('AdvancedParameters', () => {

await user.clear(numberInput);
await user.type(numberInput, '10');
await user.tab();

await waitFor(() => {
expect(Number(numberInput.value)).toBeLessThanOrEqual(6);
Expand Down
Loading