115 lines
3.1 KiB
Go
115 lines
3.1 KiB
Go
package vllm
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"time"
|
|
|
|
"iop/apps/node/internal/runtime"
|
|
)
|
|
|
|
// Capabilities probes the provider and returns the adapter's advertised
|
|
// capabilities including available targets and concurrency limits.
|
|
func (v *Vllm) Capabilities(ctx context.Context) (runtime.Capabilities, error) {
|
|
probeCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
|
|
defer cancel()
|
|
targets, err := v.fetchTargets(probeCtx)
|
|
status := runtime.ProviderStatusAvailable
|
|
if err != nil {
|
|
status = runtime.ProviderStatusUnavailable
|
|
}
|
|
return runtime.Capabilities{
|
|
AdapterName: Name,
|
|
InstanceKey: v.instanceName,
|
|
Targets: targets,
|
|
MaxConcurrency: effectiveCapacity(v.capacity, 8),
|
|
MaxQueue: v.maxQueue,
|
|
QueueTimeoutMS: v.queueTimeoutMS,
|
|
RequestTimeoutMS: v.requestTimeoutMS,
|
|
ProviderStatus: runtime.NormalizeProviderStatus(status),
|
|
}, nil
|
|
}
|
|
|
|
// ProbeProvider checks endpoint availability and target presence using the
|
|
// vLLM/SGLang /v1/models endpoint, returning only the public baseline status
|
|
// values.
|
|
func (v *Vllm) ProbeProvider(ctx context.Context, target string) (runtime.ProviderProbeResult, error) {
|
|
probeCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
|
|
defer cancel()
|
|
|
|
targets, err := v.fetchTargets(probeCtx)
|
|
|
|
result := runtime.ProviderProbeResult{
|
|
AdapterName: Name,
|
|
InstanceKey: v.instanceName,
|
|
Target: target,
|
|
}
|
|
|
|
if err != nil {
|
|
result.Status = runtime.NormalizeProviderStatus(runtime.ProviderStatusUnavailable)
|
|
result.Detail = err.Error()
|
|
return result, nil
|
|
}
|
|
|
|
result.Targets = targets
|
|
|
|
if target == "" {
|
|
result.Status = runtime.NormalizeProviderStatus(runtime.ProviderStatusAvailable)
|
|
return result, nil
|
|
}
|
|
|
|
found := false
|
|
for _, t := range targets {
|
|
if t == target {
|
|
found = true
|
|
break
|
|
}
|
|
}
|
|
|
|
if found {
|
|
result.Status = runtime.NormalizeProviderStatus(runtime.ProviderStatusAvailable)
|
|
} else {
|
|
result.Status = runtime.NormalizeProviderStatus(runtime.ProviderStatusUnavailable)
|
|
result.Detail = fmt.Sprintf("target model %q not found in provider models", target)
|
|
}
|
|
|
|
return result, nil
|
|
}
|
|
|
|
func (v *Vllm) fetchTargets(ctx context.Context) ([]string, error) {
|
|
if v.endpoint == "" {
|
|
return nil, fmt.Errorf("vllm adapter: endpoint is required")
|
|
}
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, joinURL(v.endpoint, "/v1/models"), nil)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("build request: %w", err)
|
|
}
|
|
resp, err := v.client.Do(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("request failed: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
return nil, fmt.Errorf("status code %d", resp.StatusCode)
|
|
}
|
|
var models modelsResponse
|
|
if err := json.NewDecoder(resp.Body).Decode(&models); err != nil {
|
|
return nil, fmt.Errorf("decode response: %w", err)
|
|
}
|
|
out := make([]string, 0, len(models.Data))
|
|
for _, model := range models.Data {
|
|
if model.ID != "" {
|
|
out = append(out, model.ID)
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func effectiveCapacity(capacity, defaultVal int) int {
|
|
if capacity <= 0 {
|
|
return defaultVal
|
|
}
|
|
return capacity
|
|
}
|