File size: 8,055 Bytes
9b715b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
"""Optional LangSmith tracing and structured logging for the UI Space.

Mirrors the inference Space's src/tracing.py: everything here degrades to a
no-op. If `langsmith` is not installed, or LANGSMITH_TRACING is not "true",
the decorator becomes an identity function and the helpers do nothing.

    export LANGSMITH_TRACING=true
    export LANGSMITH_API_KEY=<key>
    export LANGSMITH_PROJECT=diffusiondb-sd15-lora-inference

This is the only place that can see the actual visitor. The inference Space
is called server-to-server from here, so from its side every request appears
to come from this Space; the caller's address exists only in this process.

Visitor IP and IP-derived location are personal data in the EU/UK and several
US states. LANGSMITH_LOG_IP=off drops the raw address while keeping the
coarse city/region/country, which is usually enough for analytics.
"""

from __future__ import annotations

import logging
import os
import time
from typing import Any, Callable

logger = logging.getLogger("diffusiondb.ui")

try:
    from src import geo
except ImportError:  # written separately; treat as simply unavailable
    geo = None


def _env_flag(name: str, default: str = "") -> bool:
    return os.environ.get(name, default).strip().lower() in {"1", "true", "yes", "on"}


TRACING_ENABLED = _env_flag("LANGSMITH_TRACING")
LOG_IP = os.environ.get("LANGSMITH_LOG_IP", "on").strip().lower() not in {
    "0", "false", "no", "off",
}

_langsmith: Any = None
if TRACING_ENABLED:
    try:
        import langsmith as _langsmith  # noqa: F401
    except ImportError:
        logger.warning(
            "LANGSMITH_TRACING is set but the langsmith package is not "
            "installed; running untraced. pip install langsmith"
        )
        _langsmith = None
    else:
        if not os.environ.get("LANGSMITH_API_KEY"):
            logger.warning("LANGSMITH_TRACING is set but LANGSMITH_API_KEY is empty.")

TRACING_ACTIVE = _langsmith is not None


def traceable(**decorator_kwargs: Any) -> Callable:
    """`langsmith.traceable` when tracing is on, identity decorator otherwise."""
    if not TRACING_ACTIVE:
        def passthrough(func: Callable) -> Callable:
            return func
        return passthrough
    return _langsmith.traceable(**decorator_kwargs)


def add_trace_metadata(**fields: Any) -> None:
    """Attach key/values to the run currently in scope. No-op when untraced."""
    if not TRACING_ACTIVE or not fields:
        return
    try:
        run = _langsmith.get_current_run_tree()
        if run is not None:
            run.extra.setdefault("metadata", {}).update(fields)
    except Exception:  # never let telemetry break a generation
        logger.debug("could not attach trace metadata", exc_info=True)


def _header(request: Any, name: str) -> str:
    """Case-insensitive header read that tolerates any request shape."""
    try:
        headers = getattr(request, "headers", None) or {}
        getter = getattr(headers, "get", None)
        if getter is not None:
            value = getter(name) or getter(name.title())
            if value:
                return str(value)
        lowered = {str(k).lower(): v for k, v in dict(headers).items()}
        return str(lowered.get(name.lower(), "") or "")
    except Exception:
        return ""


def caller_ip(request: Any) -> str:
    """The visitor's address, not the proxy's.

    Spaces sit behind a reverse proxy, so request.client.host is always the
    proxy. The original address is the first entry of x-forwarded-for; the
    rest of that list is the proxy chain.
    """
    forwarded = _header(request, "x-forwarded-for")
    if forwarded:
        return forwarded.split(",")[0].strip()
    real = _header(request, "x-real-ip")
    if real:
        return real.strip()
    try:
        return str(getattr(getattr(request, "client", None), "host", "") or "")
    except Exception:
        return ""


def describe_caller(request: Any) -> dict:
    """Flat dict of what is known about the visitor. Never raises."""
    if request is None:
        return {}
    try:
        info: dict[str, Any] = {}
        ip = caller_ip(request)

        if ip and LOG_IP:
            info["ip"] = ip
        if geo is not None and ip:
            # lookup() returns {} for private ranges and on any failure.
            info.update(geo.lookup(ip))

        user_agent = _header(request, "user-agent")
        if user_agent:
            info["user_agent"] = user_agent[:200]
        language = _header(request, "accept-language")
        if language:
            info["language"] = language.split(",")[0].strip()
        session = getattr(request, "session_hash", None)
        if session:
            info["session"] = str(session)
        return info
    except Exception:
        logger.debug("could not describe caller", exc_info=True)
        return {}


def _trace_inputs(inputs: dict) -> dict:
    """Drop the unserialisable client and the raw request object.

    The request carries every header; only the fields chosen by
    describe_caller belong in a trace, and they go on as metadata instead.
    """
    return {
        key: value for key, value in inputs.items()
        if key not in {"client", "request"}
    }


def _trace_outputs(result: Any) -> dict:
    """Record the seed, never the image.

    LangSmith calls this with None when the wrapped call raised, so a bare
    unpack here would log a processing error on top of every failure.
    """
    if result is None:
        return {}
    try:
        _, seed = result
        return {"seed": seed}
    except Exception:
        return {}


def distributed_headers() -> dict:
    """HTTP headers that let the backend join this trace.

    `langsmith-trace` carries the parent run id, `baggage` carries this run's
    metadata -- which is why the caller's city/region/country end up on the
    backend's own run as well as this one, rather than only here.

    Must be called AFTER add_trace_metadata, or baggage ships without it.
    """
    if not TRACING_ACTIVE:
        return {}
    try:
        run = _langsmith.get_current_run_tree()
        return dict(run.to_headers()) if run is not None else {}
    except Exception:
        logger.debug("could not build distributed trace headers", exc_info=True)
        return {}


@traceable(
    run_type="chain",
    name="ui_generate",
    process_inputs=_trace_inputs,
    process_outputs=_trace_outputs,
)
def traced_predict(client, *, prompt, negative, steps, guidance, lora_scale,
                   seed, request=None):
    """Call the inference Space, recording the request and who made it."""
    caller = describe_caller(request)
    add_trace_metadata(**caller)

    args = (prompt, negative, steps, guidance, lora_scale, seed)
    headers = distributed_headers()

    started = time.perf_counter()
    try:
        image, used = client.predict(*args, api_name="/generate",
                                     headers=headers or None)
    except TypeError:
        # Older gradio_client has no per-call headers parameter. Losing the
        # linkage is not worth losing the generation over -- the two runs
        # just land as separate traces.
        logger.debug("gradio_client does not accept per-call headers")
        image, used = client.predict(*args, api_name="/generate")
    duration_ms = (time.perf_counter() - started) * 1000

    logger.info(
        "ui request seed=%s steps=%s lora_scale=%s ms=%.0f %s",
        used, steps, lora_scale, duration_ms,
        " ".join(f"{key}={value}" for key, value in sorted(caller.items())),
    )
    return image, used


def configure_logging(level: str | None = None) -> None:
    """Single-line structured logs. Safe to call more than once."""
    if logger.handlers:
        return
    handler = logging.StreamHandler()
    handler.setFormatter(
        logging.Formatter("%(asctime)s %(levelname)s %(name)s %(message)s")
    )
    logger.addHandler(handler)
    logger.setLevel(level or os.environ.get("LOG_LEVEL", "INFO").upper())
    logger.propagate = False