feat(osi): 打通 Go 侧真实请求验收(T-005)

This commit is contained in:
ila
2026-07-06 23:03:14 +08:00
parent 2a77a4dbc1
commit def95efb1c
12 changed files with 411 additions and 54 deletions
+111 -15
View File
@@ -1,6 +1,7 @@
package osi
import (
"bufio"
"bytes"
"context"
"encoding/json"
@@ -9,6 +10,7 @@ import (
"net"
"net/http"
"net/url"
"sort"
"strings"
"time"
@@ -21,7 +23,9 @@ type TransportConfig struct {
}
type Transport struct {
client *http.Client
client *http.Client
timeout time.Duration
dialer proxy.Dialer
}
func NewTransport(config TransportConfig) (*Transport, error) {
@@ -30,21 +34,20 @@ func NewTransport(config TransportConfig) (*Transport, error) {
timeout = 20 * time.Second
}
dialer := proxy.Dialer(proxy.Direct)
roundTripper := http.DefaultTransport.(*http.Transport).Clone()
if strings.TrimSpace(config.Socks5Proxy) != "" {
dialer, err := socks5Dialer(config.Socks5Proxy)
proxyDialer, err := socks5Dialer(config.Socks5Proxy)
if err != nil {
return nil, err
}
dialer = proxyDialer
roundTripper.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) {
if contextDialer, ok := dialer.(proxy.ContextDialer); ok {
return contextDialer.DialContext(ctx, network, address)
}
return dialer.Dial(network, address)
return dialWithContext(ctx, proxyDialer, network, address)
}
}
return &Transport{client: &http.Client{Timeout: timeout, Transport: roundTripper}}, nil
return &Transport{client: &http.Client{Timeout: timeout, Transport: roundTripper}, timeout: timeout, dialer: dialer}, nil
}
func (t *Transport) PostJSON(ctx context.Context, targetURL string, payload any, headers map[string]string) (int, []byte, error) {
@@ -53,18 +56,35 @@ func (t *Transport) PostJSON(ctx context.Context, targetURL string, payload any,
return 0, nil, fmt.Errorf("marshal json payload: %w", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
target, err := url.Parse(targetURL)
if err != nil {
return 0, nil, fmt.Errorf("build request: %w", err)
return 0, nil, fmt.Errorf("parse target url: %w", err)
}
for name, value := range headers {
req.Header.Set(name, value)
}
if req.Header.Get("Content-Type") == "" {
req.Header.Set("Content-Type", "application/json")
if target.Scheme != "http" {
return 0, nil, fmt.Errorf("unsupported target scheme %q", target.Scheme)
}
resp, err := t.client.Do(req)
reqCtx := ctx
cancel := func() {}
if _, ok := ctx.Deadline(); !ok && t.timeout > 0 {
reqCtx, cancel = context.WithTimeout(ctx, t.timeout)
}
defer cancel()
conn, err := dialWithContext(reqCtx, t.dialer, "tcp", target.Host)
if err != nil {
return 0, nil, err
}
defer conn.Close()
if t.timeout > 0 {
_ = conn.SetDeadline(time.Now().Add(t.timeout))
}
if err := writeJSONRequest(conn, target, body, headers); err != nil {
return 0, nil, err
}
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
if err != nil {
return 0, nil, err
}
@@ -77,6 +97,82 @@ func (t *Transport) PostJSON(ctx context.Context, targetURL string, payload any,
return resp.StatusCode, raw, nil
}
func writeJSONRequest(w io.Writer, target *url.URL, body []byte, headers map[string]string) error {
path := target.RequestURI()
if path == "" {
path = "/"
}
var buf bytes.Buffer
fmt.Fprintf(&buf, "POST %s HTTP/1.1\r\n", path)
writeHeader(&buf, "Host", target.Host)
writeHeader(&buf, "User-Agent", headerValue(headers, "User-Agent", "python-requests/2.32.4"))
writeHeader(&buf, "Accept-Encoding", headerValue(headers, "Accept-Encoding", "gzip, deflate, br"))
writeHeader(&buf, "Accept", headerValue(headers, "Accept", "*/*"))
writeHeader(&buf, "Connection", headerValue(headers, "Connection", "keep-alive"))
writeHeader(&buf, "Content-Length", fmt.Sprintf("%d", len(body)))
writeHeader(&buf, "Content-Type", headerValue(headers, "Content-Type", "application/json"))
written := map[string]bool{
"host": true, "user-agent": true, "accept-encoding": true, "accept": true,
"connection": true, "content-length": true, "content-type": true,
}
for _, name := range []string{"orgCode", "deviceSN", "ts", "userName", "password"} {
if value, ok := headers[name]; ok {
writeHeader(&buf, name, value)
written[strings.ToLower(name)] = true
}
}
var rest []string
for name := range headers {
if !written[strings.ToLower(name)] {
rest = append(rest, name)
}
}
sort.Strings(rest)
for _, name := range rest {
writeHeader(&buf, name, headers[name])
}
buf.WriteString("\r\n")
buf.Write(body)
_, err := w.Write(buf.Bytes())
return err
}
func writeHeader(buf *bytes.Buffer, name, value string) {
fmt.Fprintf(buf, "%s: %s\r\n", name, value)
}
func headerValue(headers map[string]string, name, fallback string) string {
if value, ok := headers[name]; ok {
return value
}
return fallback
}
func dialWithContext(ctx context.Context, dialer proxy.Dialer, network, address string) (net.Conn, error) {
if contextDialer, ok := dialer.(proxy.ContextDialer); ok {
return contextDialer.DialContext(ctx, network, address)
}
type result struct {
conn net.Conn
err error
}
ch := make(chan result, 1)
go func() {
conn, err := dialer.Dial(network, address)
ch <- result{conn: conn, err: err}
}()
select {
case <-ctx.Done():
return nil, ctx.Err()
case result := <-ch:
return result.conn, result.err
}
}
func socks5Dialer(rawProxy string) (proxy.Dialer, error) {
proxyURL, err := normalizeSocks5Proxy(rawProxy)
if err != nil {
+67
View File
@@ -3,8 +3,10 @@ package osi
import (
"context"
"encoding/json"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
@@ -77,3 +79,68 @@ func TestNewTransportRejectsInvalidSOCKS5Proxy(t *testing.T) {
t.Fatal("NewTransport accepted invalid SOCKS5 proxy")
}
}
func TestTransportPreservesHeaderNameCasing(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
defer listener.Close()
received := make(chan string, 1)
go func() {
conn, err := listener.Accept()
if err != nil {
received <- ""
return
}
defer conn.Close()
buf := make([]byte, 4096)
n, _ := conn.Read(buf)
received <- string(buf[:n])
_, _ = conn.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n{}"))
}()
transport, err := NewTransport(TransportConfig{Timeout: time.Second})
if err != nil {
t.Fatalf("NewTransport: %v", err)
}
_, _, err = transport.PostJSON(
context.Background(),
"http://"+listener.Addr().String(),
map[string]string{"serviceId": "JKDA00002"},
map[string]string{"orgCode": "org-001", "deviceSN": "", "userName": "dyytgw"},
)
if err != nil {
t.Fatalf("PostJSON: %v", err)
}
raw := <-received
for _, header := range []string{"orgCode:", "deviceSN:", "userName:"} {
if !strings.Contains(raw, "\r\n"+header) {
t.Fatalf("raw request missing exact header %q:\n%s", header, raw)
}
}
}
func TestTransportAddsRequestsCompatibleBaseHeaders(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("User-Agent") == "" {
t.Fatal("missing User-Agent")
}
if r.Header.Get("Accept") != "*/*" {
t.Fatalf("Accept = %q", r.Header.Get("Accept"))
}
if r.Header.Get("Connection") != "keep-alive" {
t.Fatalf("Connection = %q", r.Header.Get("Connection"))
}
_, _ = w.Write([]byte(`{}`))
}))
defer server.Close()
transport, err := NewTransport(TransportConfig{Timeout: time.Second})
if err != nil {
t.Fatalf("NewTransport: %v", err)
}
if _, _, err := transport.PostJSON(context.Background(), server.URL, map[string]string{"serviceId": "JKDA00002"}, nil); err != nil {
t.Fatalf("PostJSON: %v", err)
}
}